diff --git a/app/controller/add_user_route.py b/app/controller/add_user_route.py index f3df846..8edd34f 100644 --- a/app/controller/add_user_route.py +++ b/app/controller/add_user_route.py @@ -199,20 +199,8 @@ async def get_current_user(session: AsyncSessionDep, token: str): except InvalidTokenError: raise unauthorized_exception - user_model = await get_user_by_username(username, session) + user_model: UserModel | None = await UserService.query_item(session, {"username": username, "valid": "Y"}) if not user_model: raise unauthorized_exception - return PublicUser(**user_model.model_dump()) - - -# 根据用户名获取用户信息 -async def get_user_by_username(username: str, session: AsyncSessionDep): - query = ( - select(UserModel) - .where(UserModel.username == username) - .where(UserModel.valid == UserValidate.Y) - ) - result = await session.execute(query) - item_cls: UserModel | None = result.scalars().first() - return item_cls + return user_model.model_dump(mode="json") diff --git a/app/middlewares/app_middlewares.py b/app/middlewares/app_middlewares.py index a0624fe..7c0fc5c 100644 --- a/app/middlewares/app_middlewares.py +++ b/app/middlewares/app_middlewares.py @@ -10,7 +10,8 @@ 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.model.PosModel import PosModel +from app.model.UserModel import UserServiceModel 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 @@ -68,8 +69,8 @@ def add_app_middlewares(app: FastAPI): 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") + user_dict = await get_current_user(session, token) + return user_dict except InvalidTokenError: return None @@ -80,7 +81,11 @@ def add_app_middlewares(app: FastAPI): if not user_dict: raise unauthorized_exception - request.state.user = PublicUser(**user_dict) + print("user_dict", user_dict) + request.state.user = UserServiceModel(**user_dict) + user_pos_dict = user_dict.get('pos') + if user_pos_dict: + request.state.user.position = PosModel(**user_pos_dict) request.state.token = token # 如果请求的是api接口 diff --git a/app/model/BasicModel.py b/app/model/BasicModel.py index b984f93..0a83fb0 100644 --- a/app/model/BasicModel.py +++ b/app/model/BasicModel.py @@ -28,9 +28,9 @@ class BasicModel(SQLModel): # 定义datetime和date类型的JSON编码器,将其格式化为指定字符串 json_encoders = { # 若为datetime类型,格式化为“年-月-日 时:分:秒”,若为None则保持None - datetime: lambda dt: dt.strftime("%Y-%m-%d %H:%M:%S") if dt is not None else None, + datetime: lambda dt: dt if isinstance(dt, str) else dt.strftime("%Y-%m-%d %H:%M:%S") if dt is not None else None, # 若为date类型,格式化为“年-月-日”,若为None则保持None - date: lambda dt: dt.strftime("%Y-%m-%d") if dt is not None else None + date: lambda dt: dt if isinstance(dt, str) else dt.strftime("%Y-%m-%d") if dt is not None else None } # 定义模型验证器,在数据解析前(mode='before')执行,用于处理字符串格式的日期时间 diff --git a/app/utils/redis_utils.py b/app/utils/redis_utils.py index 282b308..ad184af 100644 --- a/app/utils/redis_utils.py +++ b/app/utils/redis_utils.py @@ -1,3 +1,4 @@ +import json import uuid from contextlib import asynccontextmanager from typing import Annotated, AsyncContextManager @@ -68,16 +69,17 @@ 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) + # default_value_getter返回的字段可能嵌套多层对象,这里改成用json字符串缓存 + json_string = await redis_client.get(key) + exists = bool(json_string) if exists: - return mapping + return json.loads(json_string) 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) + await redis_client.set(key, json.dumps(new_value, ensure_ascii=False)) return value