perf: CTE任务树加载器 + MOM跨库查询缓存
消除两个核心N+1性能瓶颈: 1. CTE任务树加载器 (task_tree_loader.py) - PostgreSQL Recursive CTE一次性加载完整任务树 - 无论树深度多大,仅2条SQL(CTE + records selectinload) - set_committed_value安全注入,避免Session脏数据 - 修复add_task_record双重加载问题 - 移除get_all_tasks中冗余的selectinload(child_tasks) 2. MOM跨库查询缓存 (mom_cache.py) - 零依赖TTL内存缓存(threading.RLock + time.monotonic) - 参数化ANY(:user_ids)替代OR拼接LIKE(防注入) - get_all_products中3次调用共享缓存,2h TTL内零跨库查询
This commit is contained in:
@ -46,35 +46,10 @@ def _task_to_response(task: Task) -> TaskResponse:
|
||||
|
||||
|
||||
async def _load_task_tree(db: AsyncSession, product_id: uuid.UUID) -> list[TaskResponse]:
|
||||
"""递归加载产品下的完整任务树"""
|
||||
# 先取顶层任务
|
||||
result = await db.execute(
|
||||
select(Task)
|
||||
.options(selectinload(Task.child_tasks), selectinload(Task.records))
|
||||
.where(
|
||||
Task.product_id == product_id,
|
||||
Task.parent_task_id.is_(None),
|
||||
)
|
||||
.order_by(Task.created_at)
|
||||
)
|
||||
top_tasks = result.scalars().all()
|
||||
|
||||
# 递归加载每层子任务
|
||||
async def _load_children(t: Task):
|
||||
for child in t.child_tasks:
|
||||
child_result = await db.execute(
|
||||
select(Task)
|
||||
.options(selectinload(Task.child_tasks), selectinload(Task.records))
|
||||
.where(Task.id == child.id)
|
||||
)
|
||||
refreshed = child_result.scalar_one()
|
||||
t.child_tasks[t.child_tasks.index(child)] = refreshed
|
||||
await _load_children(refreshed)
|
||||
|
||||
for task in top_tasks:
|
||||
await _load_children(task)
|
||||
|
||||
return [_task_to_response(t) for t in top_tasks]
|
||||
"""使用 PostgreSQL Recursive CTE 一次性加载产品下完整任务树(消除 N+1)"""
|
||||
from app.services.task_tree_loader import load_task_trees_by_product
|
||||
tasks = await load_task_trees_by_product(db, product_id)
|
||||
return [_task_to_response(t) for t in tasks]
|
||||
|
||||
|
||||
async def get_product_by_serial(db: AsyncSession, serial_number: str) -> ProductScanResponse:
|
||||
@ -344,32 +319,9 @@ async def update_overall_status(
|
||||
|
||||
|
||||
def _lookup_display_names(location_ids: list[str]) -> dict[str, str]:
|
||||
"""批量查询 MOM sys_user,将 username 映射为真实姓名"""
|
||||
if not location_ids:
|
||||
return {}
|
||||
from app.core.mom_database import MomSessionLocal
|
||||
from sqlalchemy import text
|
||||
db = MomSessionLocal()
|
||||
try:
|
||||
# 过滤掉特殊值
|
||||
real_ids = [uid for uid in location_ids if uid and uid != "virtual_warehouse"]
|
||||
if not real_ids:
|
||||
return {}
|
||||
# 用 LIKE 模糊匹配批量查出
|
||||
conditions = " OR ".join([f"username LIKE '%/{uid}'" for uid in real_ids])
|
||||
result = db.execute(
|
||||
text(f"SELECT username, SPLIT_PART(username, '/', 1) as display_name FROM sys_user WHERE {conditions}")
|
||||
)
|
||||
mapping = {}
|
||||
for row in result:
|
||||
full_username = row[0]
|
||||
display_name = row[1]
|
||||
# 从 full_username 末尾提取短用户名: "张三/zhangsan01" → "zhangsan01"
|
||||
short = full_username.split("/")[-1] if "/" in full_username else full_username
|
||||
mapping[short] = display_name
|
||||
return mapping
|
||||
finally:
|
||||
db.close()
|
||||
"""批量查询 MOM sys_user,将 username 映射为真实姓名(带 2h TTL 缓存)"""
|
||||
from app.services.mom_cache import get_display_names
|
||||
return get_display_names(location_ids)
|
||||
|
||||
|
||||
async def get_all_products(
|
||||
|
||||
Reference in New Issue
Block a user