feat: 修复接口认证时查询用户信息没有缓存用户的职位信息的问题
This commit is contained in:
@@ -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
|
|
||||||
|
|||||||
@@ -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接口
|
||||||
|
|||||||
@@ -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')执行,用于处理字符串格式的日期时间
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user