Files
track/backend/app/services/mom_cache.py
duxingchen cc199081f9 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内零跨库查询
2026-08-12 12:03:02 +08:00

155 lines
5.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
MOM 跨库查询缓存模块 — 使用本地 TTL 缓存消除冗余跨库请求
解决的问题:
1. _lookup_display_names 在 get_all_products 中被调用 3 次,每次都打开/关闭
MOM 数据库连接150 条产品的列表页 = 3 根管线查询。
2. 同一批 username 在短时间内(用户翻页、多人同时访问)被反复查询。
3. 旧实现用 OR 拼接 LIKE 条件,存在注入风险。
方案python -m 内置模块(零依赖)实现线程安全 TTL 缓存 + 参数化 ANY 查询。
TTL: 2 小时(人员姓名不会频繁变动,可调)。
"""
from __future__ import annotations
import threading
import time
from app.core.mom_database import MomSessionLocal
# ============================================================
# 零依赖 TTL 缓存(线程安全)
# ============================================================
class _TTLCache:
"""线程安全的内存 TTL 缓存,用于 MOM 只读查询结果"""
def __init__(self, ttl_seconds: int = 7200) -> None:
self._store: dict[str, str] = {}
self._expiry: dict[str, float] = {}
self._ttl = ttl_seconds
self._lock = threading.RLock()
def get_many(self, keys: list[str]) -> tuple[dict[str, str], list[str]]:
"""
批量获取 → (命中字典, 未命中 key 列表)。
内部自动清理过期条目。
"""
hits: dict[str, str] = {}
missed: list[str] = []
now = time.monotonic()
with self._lock:
for k in keys:
exp = self._expiry.get(k)
if exp is not None and now < exp:
hits[k] = self._store[k]
else:
missed.append(k)
# 清理过期残留
if k in self._store:
del self._store[k]
del self._expiry[k]
return hits, missed
def set_many(self, mapping: dict[str, str]) -> None:
"""批量写入,所有 key 共享同一过期时间"""
expiry = time.monotonic() + self._ttl
with self._lock:
for k, v in mapping.items():
self._store[k] = v
self._expiry[k] = expiry
# ============================================================
# 全局缓存实例2h TTL
# ============================================================
_user_name_cache = _TTLCache(ttl_seconds=7200)
# ============================================================
# 公开 API
# ============================================================
def get_display_names(user_ids: list[str]) -> dict[str, str]:
"""
批量查询 MOM sys_user将 username 映射为真实姓名(带 2h TTL 缓存)。
缓存穿透流程:
1. 去重 → 从缓存批量读取
2. 计算 miss 差集
3. miss 非空时,用参数化 ANY(:user_ids) 查 MOM1 条 SQL
4. 写回缓存
5. 合并 hits + fresh 返回
参数:
user_ids: 短用户名列表,如 ["zhangsan01", "lisi02"]
返回:
{"zhangsan01": "张三", "lisi02": "李四"}
不存在的 key 不会出现在返回字典中。
SQL 安全:
使用 SPLIT_PART(username, '/', 2) = ANY(:user_ids) 参数化查询,
杜绝旧实现中 OR 拼接 LIKE 的注入风险。
"""
if not user_ids:
return {}
# 过滤特殊值 + 去重保序
seen: set[str] = set()
real_ids: list[str] = []
for uid in user_ids:
if uid and uid != "virtual_warehouse" and uid not in seen:
seen.add(uid)
real_ids.append(uid)
if not real_ids:
return {}
# ── Step 1: 批量查缓存 ──
hits, missed = _user_name_cache.get_many(real_ids)
# ── Step 2: 仅对 miss 查 MOM ──
if missed:
db = MomSessionLocal()
try:
from sqlalchemy import text
# 参数化 ANY 查询 — 安全防注入
# SPLIT_PART('张三/zhangsan01', '/', 2) = 'zhangsan01'
# OR username = ANY(...) 兜底无斜杠的用户名(如 admin
sql = text("""
SELECT username,
SPLIT_PART(username, '/', 1) AS display_name
FROM sys_user
WHERE SPLIT_PART(username, '/', 2) = ANY(:user_ids)
OR username = ANY(:user_ids)
""")
result = db.execute(sql, {"user_ids": missed})
rows = result.fetchall()
finally:
db.close()
# ── Step 3: 解析结果 + 写回缓存 ──
fresh: dict[str, str] = {}
for row in rows:
full_username: str = row[0]
display_name: str = row[1]
# "张三/zhangsan01" → short="zhangsan01"
short = full_username.split("/")[-1] if "/" in full_username else full_username
fresh[short] = display_name
if fresh:
_user_name_cache.set_many(fresh)
# ── Step 4: 合并 ──
hits.update(fresh)
return hits