feat: 使用redis来缓存用户信息

This commit is contained in:
martsforever
2025-09-27 21:06:15 +08:00
parent 97515ecfe7
commit 4311f521a7
3 changed files with 63 additions and 15 deletions
+19 -8
View File
@@ -10,9 +10,11 @@ from starlette.responses import JSONResponse
from app.config.env import env from app.config.env import env
from app.controller.add_user_route import unauthorized_exception, get_current_user from app.controller.add_user_route import unauthorized_exception, get_current_user
from app.model.ApiSecretModel import ApiSecretService from app.model.ApiSecretModel import ApiSecretService
from app.model.UserModel import PublicUser
from app.utils.CrpyUtils import TokenInfo, CryptUtils from app.utils.CrpyUtils import TokenInfo, CryptUtils
from app.utils.api_secret_utils import api_secret_utils, ApiSecretStatus from app.utils.api_secret_utils import api_secret_utils, ApiSecretStatus
from app.utils.db_utils import async_session from app.utils.db_utils import async_session
from app.utils.redis_utils import get_redis_cache
def add_app_middlewares(app: FastAPI): def add_app_middlewares(app: FastAPI):
@@ -63,14 +65,23 @@ def add_app_middlewares(app: FastAPI):
except InvalidTokenError: except InvalidTokenError:
raise unauthorized_exception raise unauthorized_exception
# 根据token信息获取用户信息 async def get_default_cache_value():
async with async_session() as session: async with async_session() as session:
try: try:
public_user = await get_current_user(session, token) public_user = await get_current_user(session, token)
request.state.user = public_user return public_user.model_dump(mode="json")
request.state.token = token except InvalidTokenError:
except InvalidTokenError: return None
raise unauthorized_exception
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接口
# api开头的接口需要额外验证秘钥 # api开头的接口需要额外验证秘钥
+19 -5
View File
@@ -8,6 +8,7 @@ from sqlmodel import Field, Relationship, select
from app.model.BasicModel import BasicModel from app.model.BasicModel import BasicModel
from app.model.PosModel import PosModel from app.model.PosModel import PosModel
from app.utils.create_module_service import create_model_service from app.utils.create_module_service import create_model_service
from app.utils.redis_utils import remove_redis_cache
class UserValidate(str, Enum): class UserValidate(str, Enum):
@@ -93,11 +94,24 @@ class UserServiceModel(PublicUser, table=True):
pass 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( UserService = create_model_service(
Cls=UserServiceModel, Cls=UserServiceModel,
custom_query=(lambda: select(UserServiceModel) custom_query=(
.options( lambda: select(UserServiceModel)
selectinload(UserServiceModel.position). .options(
selectinload(PosModel.organization) selectinload(UserServiceModel.position).
)) selectinload(PosModel.organization)
)),
before_update=before_update,
before_delete=before_delete
) )
+25 -2
View File
@@ -1,6 +1,6 @@
import uuid import uuid
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from typing import Annotated from typing import Annotated, AsyncContextManager
import redis.asyncio as redis import redis.asyncio as redis
from fastapi import Depends from fastapi import Depends
@@ -38,7 +38,7 @@ class RedisUtils():
# 获取一个Redis连接客户端 # 获取一个Redis连接客户端
@asynccontextmanager @asynccontextmanager
async def get_redis_connection(self): async def get_redis_connection(self) -> AsyncContextManager[redis.Redis]:
if not self.redis_pool: if not self.redis_pool:
raise Exception("Redis 连接池未初始化") raise Exception("Redis 连接池未初始化")
redis_client = redis.Redis(connection_pool=self.redis_pool) 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)] 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 # 返回是否删除成功