58 lines
1.9 KiB
Python
58 lines
1.9 KiB
Python
"""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
|