perf: support keyed create locks and use user_id lock for UserConsole

This commit is contained in:
HibiKier
2025-12-23 17:19:01 +08:00
parent ed23ad319a
commit a8251165fa
2 changed files with 50 additions and 38 deletions
+7 -22
View File
@@ -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
+43 -16
View File
@@ -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():