From a8251165fa15d83f1ec12627b34754f2a284b228 Mon Sep 17 00:00:00 2001 From: HibiKier <775757368@qq.com> Date: Tue, 23 Dec 2025 17:19:01 +0800 Subject: [PATCH] perf: support keyed create locks and use user_id lock for UserConsole --- zhenxun/models/user_console.py | 29 +++-------- zhenxun/services/db_context/base_model.py | 59 +++++++++++++++++------ 2 files changed, 50 insertions(+), 38 deletions(-) diff --git a/zhenxun/models/user_console.py b/zhenxun/models/user_console.py index 096a0721..c2b4dce9 100644 --- a/zhenxun/models/user_console.py +++ b/zhenxun/models/user_console.py @@ -1,4 +1,3 @@ -import asyncio from typing import ClassVar from tortoise import fields @@ -7,7 +6,7 @@ from tortoise.exceptions import IntegrityError from zhenxun.models.goods_info import GoodsInfo from zhenxun.services.cache import CacheRoot from zhenxun.services.db_context import Model -from zhenxun.utils.enum import CacheType, GoldHandle +from zhenxun.utils.enum import CacheType, DbLockType, GoldHandle from zhenxun.utils.exception import GoodsNotFound, InsufficientGold from .user_gold_log import UserGoldLog @@ -40,32 +39,18 @@ class UserConsole(Model): """缓存类型""" cache_key_field = "user_id" """缓存键字段""" - _create_locks: ClassVar[dict[str, asyncio.Lock]] = {} + lock_fields: ClassVar[dict[DbLockType, str]] = {DbLockType.CREATE: "user_id"} @classmethod async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole": if user := await cls.get_or_none(user_id=user_id): return user - lock = cls._create_locks.setdefault(user_id, asyncio.Lock()) - async with lock: - # 双重检查,避免锁等待后重复创建 - if user := await cls.get_or_none(user_id=user_id): - return user - - for _ in range(3): - try: - uid = await cls.get_new_uid() - return await cls.create( - user_id=user_id, - uid=uid, - platform=platform, - ) - except IntegrityError: - # 可能是 uid 或 user_id 竞争,短暂等待再重试 - await asyncio.sleep(0.05) - - # 多次重试仍失败,按 user_id 再查一遍兜底 + try: + uid = await cls.get_new_uid() + return await cls.create(user_id=user_id, uid=uid, platform=platform) + except IntegrityError: + # 并发竞争下,依赖 user_id 唯一约束兜底查询 return await cls.get(user_id=user_id) @classmethod diff --git a/zhenxun/services/db_context/base_model.py b/zhenxun/services/db_context/base_model.py index ff642258..94175aa9 100644 --- a/zhenxun/services/db_context/base_model.py +++ b/zhenxun/services/db_context/base_model.py @@ -22,8 +22,11 @@ class Model(TortoiseModel): 增强的ORM基类,解决锁嵌套问题 """ - sem_data: ClassVar[dict[str, dict[str, asyncio.Semaphore]]] = {} - _current_locks: ClassVar[dict[int, DbLockType]] = {} # 跟踪当前协程持有的锁 + # sem_data[cls_name][lock_type] 可以是 Semaphore(全局) + # 或 dict[key, Semaphore](按键) + sem_data: ClassVar[dict[str, dict[str, Any]]] = {} + # 跟踪当前协程持有的锁 (lock_type, lock_key) + _current_locks: ClassVar[dict[int, tuple[DbLockType, Any | None]]] = {} def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) @@ -77,32 +80,46 @@ class Model(TortoiseModel): return None @classmethod - def get_semaphore(cls, lock_type: DbLockType): + def get_semaphore(cls, lock_type: DbLockType, lock_key: Any | None = None): enable_lock = getattr(cls, "enable_lock", None) if not enable_lock or lock_type not in enable_lock: return None - if cls.__name__ not in cls.sem_data: - cls.sem_data[cls.__name__] = {} - if lock_type not in cls.sem_data[cls.__name__]: - cls.sem_data[cls.__name__][lock_type] = asyncio.Semaphore(1) - return cls.sem_data[cls.__name__][lock_type] + cls_sem = cls.sem_data.setdefault(cls.__name__, {}) + # 是否配置了按字段的锁 + lock_fields: dict[DbLockType, str] = getattr(cls, "lock_fields", {}) or {} + if lock_type in lock_fields and lock_key is not None: + keyed = cls_sem.setdefault(lock_type, {}) + if not isinstance(keyed, dict): + # 兼容历史数据,重置为按键字典 + keyed = {} + cls_sem[lock_type] = keyed + if lock_key not in keyed: + keyed[lock_key] = asyncio.Semaphore(1) + return keyed[lock_key] + + # 默认全局锁 + sem = cls_sem.get(lock_type) + if not isinstance(sem, asyncio.Semaphore): + sem = asyncio.Semaphore(1) + cls_sem[lock_type] = sem + return sem @classmethod - def _require_lock(cls, lock_type: DbLockType) -> bool: + def _require_lock(cls, lock_type: DbLockType, lock_key: Any | None) -> bool: """检查是否需要真正加锁""" task_id = id(asyncio.current_task()) - return cls._current_locks.get(task_id) != lock_type + return cls._current_locks.get(task_id) != (lock_type, lock_key) @classmethod @contextlib.asynccontextmanager - async def _lock_context(cls, lock_type: DbLockType): + async def _lock_context(cls, lock_type: DbLockType, lock_key: Any | None = None): """带重入检查的锁上下文""" task_id = id(asyncio.current_task()) - need_lock = cls._require_lock(lock_type) + need_lock = cls._require_lock(lock_type, lock_key) - if need_lock and (sem := cls.get_semaphore(lock_type)): - cls._current_locks[task_id] = lock_type + if need_lock and (sem := cls.get_semaphore(lock_type, lock_key)): + cls._current_locks[task_id] = (lock_type, lock_key) async with sem: yield cls._current_locks.pop(task_id, None) @@ -114,7 +131,12 @@ class Model(TortoiseModel): cls, using_db: BaseDBAsyncClient | None = None, **kwargs: Any ) -> Self: """创建数据(使用CREATE锁)""" - async with cls._lock_context(DbLockType.CREATE): + lock_fields: dict[DbLockType, str] = getattr(cls, "lock_fields", {}) or {} + lock_key = None + if field := lock_fields.get(DbLockType.CREATE): + lock_key = kwargs.get(field) + + async with cls._lock_context(DbLockType.CREATE, lock_key): # 直接调用父类的_create方法避免触发save的锁 result = await super().create(using_db=using_db, **kwargs) if cache_type := cls.get_cache_type(): @@ -144,7 +166,12 @@ class Model(TortoiseModel): **kwargs: Any, ) -> tuple[Self, bool]: """更新或创建数据(使用UPSERT锁)""" - async with cls._lock_context(DbLockType.UPSERT): + lock_fields: dict[DbLockType, str] = getattr(cls, "lock_fields", {}) or {} + lock_key = None + if field := lock_fields.get(DbLockType.UPSERT): + lock_key = kwargs.get(field) + + async with cls._lock_context(DbLockType.UPSERT, lock_key): try: # 先尝试更新(带行锁) async with in_transaction():