From 40a81efa248f8dbcd260aca4b2449261740f01bc Mon Sep 17 00:00:00 2001 From: HibiKier <775757368@qq.com> Date: Wed, 31 Dec 2025 09:05:32 +0800 Subject: [PATCH] perf: fix lock ordering issue - acquire semaphore before checking building status --- zhenxun/services/auth_snapshot/service.py | 151 +++++++--------------- 1 file changed, 50 insertions(+), 101 deletions(-) diff --git a/zhenxun/services/auth_snapshot/service.py b/zhenxun/services/auth_snapshot/service.py index 41e12f22..91fbf1f7 100644 --- a/zhenxun/services/auth_snapshot/service.py +++ b/zhenxun/services/auth_snapshot/service.py @@ -45,9 +45,6 @@ class AuthSnapshotService: # 正在构建中的快照(防止并发重复构建) _building: ClassVar[dict[str, asyncio.Future]] = {} - # 构建锁(按 cache_key 粒度) - _build_locks: ClassVar[dict[str, asyncio.Lock]] = {} - # 全局构建并发限制(防止大量不同 key 同时构建导致 DB 过载) _build_semaphore: ClassVar[asyncio.Semaphore | None] = None @@ -100,7 +97,7 @@ class AuthSnapshotService: memory_cache = cls._get_memory_cache() - # 1. 尝试从内存缓存获取 + # 1. 尝试从内存缓存获取(最快路径) if not force_refresh: if snapshot := memory_cache.get(cache_key): return snapshot @@ -112,83 +109,72 @@ class AuthSnapshotService: if cached and isinstance(cached, dict): snapshot = AuthSnapshot.model_validate(cached) if not snapshot.is_expired(AUTH_REDIS_TTL): - # 更新内存缓存 memory_cache.set(cache_key, snapshot) return snapshot except Exception as e: - logger.debug(f"从Redis获取权限快照失败: {cache_key}", LOG_COMMAND, e=e) + logger.debug(f"从Redis获取快照失败: {cache_key}", LOG_COMMAND, e=e) - # 3. 获取或创建该 cache_key 的锁 - if cache_key not in cls._build_locks: - cls._build_locks[cache_key] = asyncio.Lock() - lock = cls._build_locks[cache_key] + # 3. 检查是否有其他协程正在构建(无需持锁) + if cache_key in cls._building: + try: + return await cls._building[cache_key] + except Exception: + pass - # 4. 使用锁保护构建过程,防止并发重复构建 - async with lock: - # 再次检查缓存(可能在等待锁的过程中已被其他协程构建) + # 4. 需要构建 - 先获取信号量(不在锁内等待) + semaphore = cls._get_build_semaphore() + try: + await asyncio.wait_for(semaphore.acquire(), timeout=BUILD_QUEUE_TIMEOUT) + except asyncio.TimeoutError: + logger.warning(f"获取信号量超时,使用默认快照: {cache_key}", LOG_COMMAND) + return AuthSnapshot(user_id=user_id, group_id=group_id, bot_id=bot_id) + + try: + # 5. 获取信号量后,再次检查缓存和构建状态 if snapshot := memory_cache.get(cache_key): return snapshot - # 检查是否正在构建中(其他协程已开始构建) if cache_key in cls._building: try: return await cls._building[cache_key] except Exception: pass - # 5. 构建新快照 - return await cls._build_and_cache(user_id, group_id, bot_id, cache_key) + # 6. 真正开始构建 + return await cls._do_build(user_id, group_id, bot_id, cache_key) + finally: + semaphore.release() @classmethod - async def _build_and_cache( + async def _do_build( cls, user_id: str, group_id: str | None, bot_id: str, cache_key: str, ) -> AuthSnapshot: - """构建并缓存快照(带全局并发限制)""" + """执行快照构建(信号量已在外部获取)""" loop = asyncio.get_running_loop() future: asyncio.Future[AuthSnapshot] = loop.create_future() cls._building[cache_key] = future - semaphore = cls._get_build_semaphore() - try: - # 尝试获取信号量(限制并发构建数量) - try: - await asyncio.wait_for(semaphore.acquire(), timeout=BUILD_QUEUE_TIMEOUT) - except asyncio.TimeoutError: - # 等待超时,返回默认快照(允许请求继续,但不保证权限数据完整) - logger.warning( - f"构建快照等待超时(并发过高),使用默认快照: {cache_key}", - LOG_COMMAND, - ) - snapshot = AuthSnapshot( - user_id=user_id, group_id=group_id, bot_id=bot_id - ) - future.set_result(snapshot) - return snapshot + # 构建快照 + snapshot = await SnapshotBuilder.build_auth_snapshot( + user_id, group_id, bot_id + ) - try: - # 构建快照 - snapshot = await SnapshotBuilder.build_auth_snapshot( - user_id, group_id, bot_id + # 存入Redis缓存(异步,不阻塞) + if cache_config.cache_mode != CacheMode.NONE: + asyncio.create_task( # noqa: RUF006 + cls._cache_to_redis(cache_key, snapshot) ) - # 存入Redis缓存(异步,不阻塞) - if cache_config.cache_mode != CacheMode.NONE: - asyncio.create_task( # noqa: RUF006 - cls._cache_to_redis(cache_key, snapshot) - ) + # 存入内存缓存 + cls._get_memory_cache().set(cache_key, snapshot) - # 存入内存缓存 - cls._get_memory_cache().set(cache_key, snapshot) - - future.set_result(snapshot) - return snapshot - finally: - semaphore.release() + future.set_result(snapshot) + return snapshot except Exception as e: future.set_exception(e) @@ -259,20 +245,9 @@ class AuthSnapshotService: """清空所有缓存""" if cls._memory_cache: cls._memory_cache.clear() - cls._build_locks.clear() + cls._building.clear() logger.info("已清空所有权限快照缓存", LOG_COMMAND) - @classmethod - def cleanup_locks(cls): - """清理未被使用的锁(可定期调用)""" - # 只保留正在使用的锁 - active_keys = set(cls._building.keys()) - keys_to_remove = [k for k in cls._build_locks if k not in active_keys] - for key in keys_to_remove: - lock = cls._build_locks.get(key) - if lock and not lock.locked(): - del cls._build_locks[key] - class PluginSnapshotService: """插件快照服务 @@ -286,9 +261,6 @@ class PluginSnapshotService: # 正在构建中的快照 _building: ClassVar[dict[str, asyncio.Future]] = {} - # 构建锁(按 cache_key 粒度) - _build_locks: ClassVar[dict[str, asyncio.Lock]] = {} - @classmethod def _get_memory_cache(cls) -> CacheDict[PluginSnapshot]: """获取内存缓存实例(懒加载)""" @@ -321,7 +293,7 @@ class PluginSnapshotService: cache_key = cls._build_cache_key(module) memory_cache = cls._get_memory_cache() - # 1. 尝试从内存缓存获取(最快) + # 1. 尝试从内存缓存获取(最快路径) if not force_refresh: if snapshot := memory_cache.get(cache_key): return snapshot @@ -338,32 +310,19 @@ class PluginSnapshotService: except Exception as e: logger.debug(f"从Redis获取插件快照失败: {module}", LOG_COMMAND, e=e) - # 3. 获取或创建该 cache_key 的锁 - if cache_key not in cls._build_locks: - cls._build_locks[cache_key] = asyncio.Lock() - lock = cls._build_locks[cache_key] + # 3. 检查是否有其他协程正在构建 + if cache_key in cls._building: + try: + return await cls._building[cache_key] + except Exception: + pass - # 4. 使用锁保护构建过程 - async with lock: - # 再次检查缓存(可能在等待锁的过程中已被其他协程构建) - if snapshot := memory_cache.get(cache_key): - return snapshot - - # 检查是否正在构建中 - if cache_key in cls._building: - try: - return await cls._building[cache_key] - except Exception: - pass - - # 5. 从数据库构建 - return await cls._build_and_cache(module, cache_key) + # 4. 从数据库构建(插件数量有限,无需信号量) + return await cls._do_build(module, cache_key) @classmethod - async def _build_and_cache( - cls, module: str, cache_key: str - ) -> PluginSnapshot | None: - """构建并缓存插件快照""" + async def _do_build(cls, module: str, cache_key: str) -> PluginSnapshot | None: + """执行插件快照构建""" loop = asyncio.get_running_loop() future: asyncio.Future[PluginSnapshot | None] = loop.create_future() cls._building[cache_key] = future @@ -372,7 +331,7 @@ class PluginSnapshotService: snapshot = await SnapshotBuilder.build_plugin_snapshot(module) if snapshot: - # 存入Redis缓存 + # 存入Redis缓存(异步) if cache_config.cache_mode != CacheMode.NONE: asyncio.create_task( # noqa: RUF006 cls._cache_to_redis(cache_key, snapshot) @@ -468,15 +427,5 @@ class PluginSnapshotService: """清空所有缓存""" if cls._memory_cache: cls._memory_cache.clear() - cls._build_locks.clear() + cls._building.clear() logger.info("已清空所有插件快照缓存", LOG_COMMAND) - - @classmethod - def cleanup_locks(cls): - """清理未被使用的锁(可定期调用)""" - active_keys = set(cls._building.keys()) - keys_to_remove = [k for k in cls._build_locks if k not in active_keys] - for key in keys_to_remove: - lock = cls._build_locks.get(key) - if lock and not lock.locked(): - del cls._build_locks[key]