From 587396eb490c6860b5c6683babd93ef68e94646f Mon Sep 17 00:00:00 2001 From: HibiKier <775757368@qq.com> Date: Tue, 23 Dec 2025 17:23:38 +0800 Subject: [PATCH] refactor: remove enable_lock attribute and enhance locking mechanism in Model class --- zhenxun/models/ban_console.py | 2 - zhenxun/models/group_console.py | 2 - zhenxun/services/__init__.py | 1 - zhenxun/services/db_context/base_model.py | 67 +++++++++++++++++------ 4 files changed, 51 insertions(+), 21 deletions(-) diff --git a/zhenxun/models/ban_console.py b/zhenxun/models/ban_console.py index bf41d9ba..672a6643 100644 --- a/zhenxun/models/ban_console.py +++ b/zhenxun/models/ban_console.py @@ -42,8 +42,6 @@ class BanConsole(Model): """缓存类型""" cache_key_field = ("user_id", "group_id") """缓存键字段""" - enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT] - """开启锁""" lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = { DbLockType.CREATE: ("user_id", "group_id"), DbLockType.UPSERT: ("user_id", "group_id"), diff --git a/zhenxun/models/group_console.py b/zhenxun/models/group_console.py index 7b346a67..8e4383e1 100644 --- a/zhenxun/models/group_console.py +++ b/zhenxun/models/group_console.py @@ -96,8 +96,6 @@ class GroupConsole(Model): """缓存类型""" cache_key_field = ("group_id", "channel_id") """缓存键字段""" - enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT] - """开启锁""" lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = { DbLockType.CREATE: ("group_id", "channel_id"), DbLockType.UPSERT: ("group_id", "channel_id"), diff --git a/zhenxun/services/__init__.py b/zhenxun/services/__init__.py index 209cc2ec..10b271d8 100644 --- a/zhenxun/services/__init__.py +++ b/zhenxun/services/__init__.py @@ -94,7 +94,6 @@ __all__ = [ async def cancel_pending_tasks(): - """仅取消 zhenxun 业务协程,避免干扰框架的关机流程。""" loop = asyncio.get_running_loop() current = asyncio.current_task(loop=loop) pending = [] diff --git a/zhenxun/services/db_context/base_model.py b/zhenxun/services/db_context/base_model.py index 63b07579..5e3bd19b 100644 --- a/zhenxun/services/db_context/base_model.py +++ b/zhenxun/services/db_context/base_model.py @@ -22,11 +22,13 @@ class Model(TortoiseModel): 增强的ORM基类,解决锁嵌套问题 """ - # sem_data[cls_name][lock_type] 可以是 Semaphore(全局) + # sem_data[cls][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]]] = {} + sem_data: ClassVar[dict[type["Model"], dict[DbLockType, Any]]] = {} + # 跟踪当前协程持有的锁集合 {(cls, lock_type, lock_key), ...} + _current_locks: ClassVar[ + dict[int, set[tuple[type["Model"], DbLockType, Any | None]]] + ] = {} def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) @@ -81,14 +83,26 @@ class Model(TortoiseModel): @classmethod 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: + """ + 获取信号量 + + 设计约定(弃用 enable_lock,仅通过 lock_fields 控制是否启用锁): + - 如果未配置 lock_fields,或其中不存在对应 lock_type,则不加锁 + - 如果 lock_fields[lock_type] 配置了按字段的锁(如 tuple[str, ...]), + 则调用处按字段值生成 lock_key,在此为不同 lock_key + 分配不同信号量,实现「按键」互斥 + - 如仅需全局锁,可在 lock_fields 中声明该 lock_type, + 且在 _lock_context 传入 lock_key=None + """ + lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {} + # 未在 lock_fields 中声明的 lock_type 不加锁 + if lock_type not in lock_fields: return None - cls_sem = cls.sem_data.setdefault(cls.__name__, {}) - # 是否配置了按字段的锁 - lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {} - if lock_type in lock_fields and lock_key is not None: + cls_sem = cls.sem_data.setdefault(cls, {}) + + # 配置了按字段的锁并且提供了具体的 lock_key 时,使用「按键」锁 + if lock_key is not None: keyed = cls_sem.setdefault(lock_type, {}) if not isinstance(keyed, dict): # 兼容历史数据,重置为按键字典 @@ -109,7 +123,13 @@ class Model(TortoiseModel): 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, lock_key) + held = cls._current_locks.get(task_id) + if not held: + return True + # 同一协程内,如果已经持有完全相同的一把锁 + # (同一模型 + 同一 lock_type + 同一 lock_key),视为重入, + # 不再重复加锁,避免自锁 + return (cls, lock_type, lock_key) not in held @classmethod @contextlib.asynccontextmanager @@ -118,13 +138,28 @@ class Model(TortoiseModel): task_id = id(asyncio.current_task()) need_lock = cls._require_lock(lock_type, lock_key) - if need_lock and (sem := cls.get_semaphore(lock_type, lock_key)): - cls._current_locks[task_id] = (lock_type, lock_key) + if not need_lock: + # 已经持有这把锁,直接透传,支持可重入 + yield + return + + sem = cls.get_semaphore(lock_type, lock_key) + if not sem: + # 对于未启用锁的场景,直接继续执行 + yield + return + + lock_id = (cls, lock_type, lock_key) + held = cls._current_locks.setdefault(task_id, set()) + held.add(lock_id) + try: async with sem: yield - cls._current_locks.pop(task_id, None) - else: - yield + finally: + # 安全移除当前锁记录 + held.discard(lock_id) + if not held: + cls._current_locks.pop(task_id, None) @classmethod async def create(