54 lines
1.5 KiB
Python
54 lines
1.5 KiB
Python
"""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"]
|