"""OEM 用户与 oem 表主键解析(全局缓存,多模块共用)。""" from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.models.oem import Oem from app.models.user import User # user_id -> oem 表主键(desktop_configs.oem_id 等同此值) _oem_id_by_user_id: dict[int, int] = {} def invalidate_oem_id_cache(user_id: int | None = None) -> None: """清除 resolve 缓存;user_id 为 None 时清空全部。""" if user_id is None: _oem_id_by_user_id.clear() else: _oem_id_by_user_id.pop(user_id, None) async def resolve_oem_id_for_user_id(db: AsyncSession, user_id: int) -> int: """ 解析 OEM 用户在业务表中的 oem_id(oem 表主键)。 若尚无 oem 记录,回退为 user_id(与 desktop_configs 历史逻辑一致)。 """ cached = _oem_id_by_user_id.get(user_id) if cached is not None: return cached oem_pk = await db.scalar(select(Oem.id).where(Oem.user_id == user_id)) resolved = int(oem_pk) if oem_pk is not None else user_id _oem_id_by_user_id[user_id] = resolved return resolved async def resolve_oem_id(db: AsyncSession, oem_user: User) -> int: """根据 User 模型解析 oem_id。""" return await resolve_oem_id_for_user_id(db, oem_user.id) async def get_or_create_oem_for_user(db: AsyncSession, user_id: int) -> Oem: """按 user_id 获取 oem 行;不存在则创建并刷新缓存。""" existing = await db.scalar(select(Oem).where(Oem.user_id == user_id)) if existing is not None: _oem_id_by_user_id[user_id] = existing.id return existing owner = await db.get(User, user_id) if owner is None: raise ValueError("关联用户不存在") oem = Oem(user_id=user_id, software_name="") db.add(oem) await db.flush() await db.refresh(oem) _oem_id_by_user_id[user_id] = oem.id return oem