feat: 使用redis来缓存用户信息
This commit is contained in:
@@ -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,15 +65,24 @@ 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
|
||||||
|
|
||||||
|
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
|
raise unauthorized_exception
|
||||||
|
|
||||||
|
request.state.user = PublicUser(**user_dict)
|
||||||
|
request.state.token = token
|
||||||
|
|
||||||
# 如果请求的是api接口
|
# 如果请求的是api接口
|
||||||
# api开头的接口需要额外验证秘钥
|
# api开头的接口需要额外验证秘钥
|
||||||
if request.url.path.startswith("/api/"):
|
if request.url.path.startswith("/api/"):
|
||||||
|
|||||||
+16
-2
@@ -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=(
|
||||||
|
lambda: select(UserServiceModel)
|
||||||
.options(
|
.options(
|
||||||
selectinload(UserServiceModel.position).
|
selectinload(UserServiceModel.position).
|
||||||
selectinload(PosModel.organization)
|
selectinload(PosModel.organization)
|
||||||
))
|
)),
|
||||||
|
before_update=before_update,
|
||||||
|
before_delete=before_delete
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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 # 返回是否删除成功
|
||||||
|
|||||||
Reference in New Issue
Block a user