Files
ai-talk-callback/app/redis_lock.py
2025-12-11 11:23:26 +08:00

334 lines
10 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 close(self):
"""断开Redis连接 (disconnect方法的别名)"""
await self.disconnect()
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()