"""Redis 会话归档分布式锁。""" from __future__ import annotations import uuid class SessionArchiveLock: """保证同一会话同一时刻只被一个归档任务处理。""" KEY = "session:archive:lock:{session_id}" DEFAULT_TTL = 120 RELEASE_SCRIPT = """ if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) end return 0 """ def __init__(self, redis, *, default_ttl: int = DEFAULT_TTL): self.redis = redis self.default_ttl = default_ttl @classmethod def key(cls, session_id: str) -> str: """生成会话级锁 Key。""" return cls.KEY.format(session_id=session_id) async def acquire(self, session_id: str, *, ttl: int | None = None) -> str | None: """尝试获取锁,成功返回随机 token,失败返回 None。""" lock_ttl = self.default_ttl if ttl is None else ttl if lock_ttl <= 0: raise ValueError("锁 TTL 必须为正整数") token = uuid.uuid4().hex acquired = await self.redis.set( self.key(session_id), token, nx=True, ex=lock_ttl ) return token if acquired else None async def release(self, session_id: str, token: str) -> bool: """仅当 token 属于当前持有者时释放锁。""" if not token: return False result = await self.redis.eval( self.RELEASE_SCRIPT, 1, self.key(session_id), token, ) return bool(result) __all__ = ["SessionArchiveLock"]