330 lines
10 KiB
Python
330 lines
10 KiB
Python
import asyncio
|
||
import uuid
|
||
from typing import Optional
|
||
from contextlib import asynccontextmanager
|
||
import aioredis
|
||
from app.config import settings
|
||
from app.logger import get_redis_logger
|
||
|
||
logger = get_redis_logger()
|
||
|
||
|
||
class AsyncRedisLock:
|
||
"""异步Redis分布式锁"""
|
||
def __init__(self, redis_pool, key: str, timeout: int = None):
|
||
self.redis_pool = redis_pool
|
||
self.key = f"lock:{key}"
|
||
self.timeout = timeout or settings.redis_lock_timeout
|
||
self.identifier = str(uuid.uuid4())
|
||
self.acquired = False
|
||
|
||
async def acquire(self, blocking: bool = False) -> bool:
|
||
"""获取分布式锁"""
|
||
logger.debug(f"🔒 尝试获取Redis锁: {self.key}")
|
||
|
||
try:
|
||
# 使用SET命令的NX和EX选项原子性地获取锁
|
||
result = await self.redis_pool.set(
|
||
self.key,
|
||
self.identifier,
|
||
expire=self.timeout,
|
||
exist=self.redis_pool.SET_IF_NOT_EXIST
|
||
)
|
||
|
||
self.acquired = result
|
||
if self.acquired:
|
||
logger.debug(f"✅ Redis锁获取成功: {self.key}")
|
||
else:
|
||
logger.debug(f"❌ Redis锁获取失败: {self.key}")
|
||
|
||
return self.acquired
|
||
except Exception as e:
|
||
logger.error(f"❌ 获取Redis锁失败: {e}")
|
||
return False
|
||
|
||
async def release(self) -> bool:
|
||
"""释放分布式锁"""
|
||
if not self.acquired:
|
||
return False
|
||
|
||
try:
|
||
# 使用Lua脚本确保只有锁的持有者才能释放锁
|
||
lua_script = """
|
||
if redis.call("GET", KEYS[1]) == ARGV[1] then
|
||
return redis.call("DEL", KEYS[1])
|
||
else
|
||
return 0
|
||
end
|
||
"""
|
||
|
||
result = await self.redis_pool.eval(
|
||
lua_script,
|
||
1,
|
||
self.key,
|
||
self.identifier
|
||
)
|
||
|
||
self.acquired = False
|
||
released = bool(result)
|
||
if released:
|
||
logger.debug(f"🔓 Redis锁释放成功: {self.key}")
|
||
else:
|
||
logger.warning(f"⚠️ Redis锁释放失败,可能已过期: {self.key}")
|
||
|
||
return released
|
||
except Exception as e:
|
||
logger.error(f"❌ 释放Redis锁失败: {e}")
|
||
return False
|
||
|
||
async def __aenter__(self):
|
||
"""异步上下文管理器入口"""
|
||
retries = 0
|
||
while retries < settings.redis_lock_max_retries:
|
||
if await self.acquire(blocking=False):
|
||
return self
|
||
await asyncio.sleep(settings.redis_lock_retry_delay)
|
||
retries += 1
|
||
|
||
raise TimeoutError(f"Failed to acquire lock {self.key} after {retries} retries")
|
||
|
||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||
"""异步上下文管理器出口"""
|
||
await self.release()
|
||
|
||
|
||
@asynccontextmanager
|
||
async def distributed_lock(redis_pool, lock_key: str, timeout: int = None):
|
||
"""分布式锁上下文管理器
|
||
|
||
Args:
|
||
redis_pool: Redis连接池
|
||
lock_key: 锁键名
|
||
timeout: 锁超时时间(秒)
|
||
|
||
Yields:
|
||
bool: 是否成功获取锁
|
||
"""
|
||
full_lock_key = f"lock:{lock_key}"
|
||
lock_timeout = timeout or settings.redis_lock_timeout
|
||
|
||
try:
|
||
# 尝试获取锁,使用setnx命令,并设置过期时间
|
||
identifier = str(uuid.uuid4())
|
||
lock_acquired = await redis_pool.set(
|
||
full_lock_key,
|
||
identifier,
|
||
expire=lock_timeout,
|
||
exist=redis_pool.SET_IF_NOT_EXIST
|
||
)
|
||
|
||
logger.debug(f"🔒 尝试获取分布式锁: {full_lock_key}, 结果: {lock_acquired}")
|
||
|
||
if lock_acquired:
|
||
try:
|
||
yield True
|
||
finally:
|
||
# 使用Lua脚本安全释放锁,确保只有锁的持有者才能释放
|
||
lua_script = """
|
||
if redis.call("GET", KEYS[1]) == ARGV[1] then
|
||
return redis.call("DEL", KEYS[1])
|
||
else
|
||
return 0
|
||
end
|
||
"""
|
||
|
||
result = await redis_pool.eval(
|
||
lua_script,
|
||
1,
|
||
full_lock_key,
|
||
identifier
|
||
)
|
||
|
||
if result:
|
||
logger.debug(f"🔓 分布式锁释放成功: {full_lock_key}")
|
||
else:
|
||
logger.warning(f"⚠️ 分布式锁释放失败,可能已过期: {full_lock_key}")
|
||
else:
|
||
yield False
|
||
|
||
except Exception as e:
|
||
logger.error(f"❌ 分布式锁操作失败: {e}")
|
||
yield False
|
||
|
||
|
||
class AsyncRedisManager:
|
||
def __init__(self):
|
||
self.redis_pool: Optional[aioredis.Redis] = None
|
||
|
||
async def connect(self) -> bool:
|
||
"""连接Redis
|
||
|
||
Returns:
|
||
bool: 连接是否成功
|
||
"""
|
||
try:
|
||
logger.info(f"🔴 正在连接Redis: {settings.redis_url}")
|
||
self.redis_pool = await aioredis.create_redis_pool(
|
||
settings.redis_url,
|
||
encoding="utf-8",
|
||
minsize=1,
|
||
maxsize=settings.redis_max_connections,
|
||
timeout=settings.redis_timeout
|
||
)
|
||
|
||
# 测试连接
|
||
result = await self.redis_pool.ping()
|
||
if result:
|
||
logger.info("✅ Redis连接成功")
|
||
return True
|
||
else:
|
||
logger.error("❌ Redis ping失败")
|
||
return False
|
||
except Exception as e:
|
||
logger.error(f"❌ Redis连接失败: {e}")
|
||
return False
|
||
|
||
async def disconnect(self):
|
||
"""断开Redis连接"""
|
||
if self.redis_pool:
|
||
self.redis_pool.close()
|
||
await self.redis_pool.wait_closed()
|
||
logger.info("🔴 Redis连接已关闭")
|
||
|
||
async def create_lock(self, key: str, timeout: int = None) -> AsyncRedisLock:
|
||
"""创建分布式锁"""
|
||
if not self.redis_pool:
|
||
raise RuntimeError("Redis client not connected")
|
||
return AsyncRedisLock(self.redis_pool, key, timeout)
|
||
|
||
async def get_redis_pool(self):
|
||
"""获取Redis连接池"""
|
||
if not self.redis_pool:
|
||
await self.connect()
|
||
return self.redis_pool
|
||
|
||
|
||
# 为了保持向后兼容,保留同步版本但标记为废弃
|
||
class RedisLock:
|
||
"""Redis分布式锁 (已废弃,请使用AsyncRedisLock)"""
|
||
def __init__(self, redis_client, key: str, timeout: int = None):
|
||
self.redis_client = redis_client
|
||
self.key = f"lock:{key}"
|
||
self.timeout = timeout or settings.redis_lock_timeout
|
||
self.identifier = str(uuid.uuid4())
|
||
self.acquired = False
|
||
|
||
def acquire(self, blocking: bool = True, timeout: int = None) -> bool:
|
||
"""获取分布式锁"""
|
||
logger.debug(f"🔒 尝试获取Redis锁: {self.key}")
|
||
|
||
lua_script = """
|
||
if redis.call("GET", KEYS[1]) == false then
|
||
return redis.call("SETEX", KEYS[1], ARGV[1], ARGV[2])
|
||
else
|
||
return false
|
||
end
|
||
"""
|
||
|
||
result = self.redis_client.eval(
|
||
lua_script,
|
||
1,
|
||
self.key,
|
||
self.timeout,
|
||
self.identifier
|
||
)
|
||
|
||
self.acquired = bool(result)
|
||
if self.acquired:
|
||
logger.debug(f"✅ Redis锁获取成功: {self.key}")
|
||
else:
|
||
logger.debug(f"❌ Redis锁获取失败: {self.key}")
|
||
|
||
return self.acquired
|
||
|
||
def release(self) -> bool:
|
||
"""释放分布式锁"""
|
||
if not self.acquired:
|
||
return False
|
||
|
||
lua_script = """
|
||
if redis.call("GET", KEYS[1]) == ARGV[1] then
|
||
return redis.call("DEL", KEYS[1])
|
||
else
|
||
return 0
|
||
end
|
||
"""
|
||
|
||
result = self.redis_client.eval(
|
||
lua_script,
|
||
1,
|
||
self.key,
|
||
self.identifier
|
||
)
|
||
|
||
self.acquired = False
|
||
return bool(result)
|
||
|
||
def __enter__(self):
|
||
"""同步上下文管理器入口"""
|
||
retries = 0
|
||
while retries < settings.redis_lock_max_retries:
|
||
if self.acquire(blocking=False):
|
||
return self
|
||
import time
|
||
time.sleep(settings.redis_lock_retry_delay)
|
||
retries += 1
|
||
|
||
raise TimeoutError(f"Failed to acquire lock {self.key} after {retries} retries")
|
||
|
||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||
"""同步上下文管理器出口"""
|
||
self.release()
|
||
|
||
|
||
class RedisManager:
|
||
def __init__(self):
|
||
self.redis_client = None
|
||
|
||
def connect(self) -> bool:
|
||
"""连接Redis (已废弃,请使用AsyncRedisManager)
|
||
|
||
Returns:
|
||
bool: 连接是否成功
|
||
"""
|
||
import redis
|
||
try:
|
||
logger.info(f"🔴 正在连接Redis: {settings.redis_url}")
|
||
self.redis_client = redis.from_url(
|
||
settings.redis_url,
|
||
encoding="utf-8",
|
||
decode_responses=True,
|
||
socket_connect_timeout=5,
|
||
socket_timeout=5,
|
||
retry_on_timeout=True
|
||
)
|
||
result = self.redis_client.ping()
|
||
if result:
|
||
logger.info("✅ Redis连接成功")
|
||
return True
|
||
else:
|
||
logger.error("❌ Redis ping失败")
|
||
return False
|
||
except Exception as e:
|
||
logger.error(f"❌ Redis连接失败: {e}")
|
||
return False
|
||
|
||
def disconnect(self):
|
||
"""断开Redis连接"""
|
||
if self.redis_client:
|
||
self.redis_client.close()
|
||
|
||
def create_lock(self, key: str, timeout: int = None) -> RedisLock:
|
||
"""创建分布式锁"""
|
||
if not self.redis_client:
|
||
raise RuntimeError("Redis client not connected")
|
||
return RedisLock(self.redis_client, key, timeout)
|
||
|
||
|
||
# 全局Redis管理器实例(保持向后兼容)
|
||
redis_manager = RedisManager()
|
||
async_redis_manager = AsyncRedisManager() |