feat: 修复接口认证时查询用户信息没有缓存用户的职位信息的问题

This commit is contained in:
martsforever
2025-09-29 13:26:39 +08:00
parent 10f606bd15
commit 498ee970f5
4 changed files with 19 additions and 24 deletions
+2 -14
View File
@@ -199,20 +199,8 @@ async def get_current_user(session: AsyncSessionDep, token: str):
except InvalidTokenError: except InvalidTokenError:
raise unauthorized_exception 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: if not user_model:
raise unauthorized_exception raise unauthorized_exception
return PublicUser(**user_model.model_dump()) return user_model.model_dump(mode="json")
# 根据用户名获取用户信息
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
+9 -4
View File
@@ -10,7 +10,8 @@ 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.model.PosModel import PosModel
from app.model.UserModel import UserServiceModel
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
@@ -68,8 +69,8 @@ def add_app_middlewares(app: FastAPI):
async def get_default_cache_value(): 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) user_dict = await get_current_user(session, token)
return public_user.model_dump(mode="json") return user_dict
except InvalidTokenError: except InvalidTokenError:
return None return None
@@ -80,7 +81,11 @@ def add_app_middlewares(app: FastAPI):
if not user_dict: if not user_dict:
raise unauthorized_exception 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 request.state.token = token
# 如果请求的是api接口 # 如果请求的是api接口
+2 -2
View File
@@ -28,9 +28,9 @@ class BasicModel(SQLModel):
# 定义datetime和date类型的JSON编码器,将其格式化为指定字符串 # 定义datetime和date类型的JSON编码器,将其格式化为指定字符串
json_encoders = { json_encoders = {
# 若为datetime类型,格式化为“年-月-日 时:分:秒”,若为None则保持None # 若为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类型,格式化为“年-月-日”,若为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')执行,用于处理字符串格式的日期时间 # 定义模型验证器,在数据解析前(mode='before')执行,用于处理字符串格式的日期时间
+6 -4
View File
@@ -1,3 +1,4 @@
import json
import uuid import uuid
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from typing import Annotated, AsyncContextManager from typing import Annotated, AsyncContextManager
@@ -68,16 +69,17 @@ RedisClientDep = Annotated[redis.Redis, Depends(get_redis_client)]
# 从redis中获取key的缓存,如果没有值则执行默认值获取函数,并保存到redis中 # 从redis中获取key的缓存,如果没有值则执行默认值获取函数,并保存到redis中
async def get_redis_cache(key: str, default_value_getter): async def get_redis_cache(key: str, default_value_getter):
async with redis_utils.get_redis_connection() as redis_client: async with redis_utils.get_redis_connection() as redis_client:
mapping = await redis_client.hgetall(key) # default_value_getter返回的字段可能嵌套多层对象,这里改成用json字符串缓存
exists = bool(mapping) json_string = await redis_client.get(key)
exists = bool(json_string)
if exists: if exists:
return mapping return json.loads(json_string)
else: else:
value = await default_value_getter() value = await default_value_getter()
print("value", value) print("value", value)
if value is not None: if value is not None:
new_value = {k: v for k, v in value.items() if v 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 return value