celery任务改为异步访问数据库

This commit is contained in:
mark.tian
2025-12-10 11:40:56 +08:00
parent d3d3502f5a
commit 572cc3113f
5 changed files with 456 additions and 213 deletions

9
2.0.0 Normal file
View File

@@ -0,0 +1,9 @@
Looking in indexes: https://mirrors.aliyun.com/pypi/simple/, https://pypi.tuna.tsinghua.edu.cn/simple
Collecting aioredis
Using cached https://pypi.tuna.tsinghua.edu.cn/packages/9b/a9/0da089c3ae7a31cbcd2dcf0214f6f571e1295d292b6139e2bac68ec081d0/aioredis-2.0.1-py3-none-any.whl (71 kB)
Collecting async-timeout (from aioredis)
Downloading https://pypi.tuna.tsinghua.edu.cn/packages/fe/ba/e2081de779ca30d473f21f5b30e0e737c438205440784c7dfc81efc2b029/async_timeout-5.0.1-py3-none-any.whl (6.2 kB)
Requirement already satisfied: typing-extensions in e:\ai\code\ai-talk-callback\.venv\lib\site-packages (from aioredis) (4.15.0)
Installing collected packages: async-timeout, aioredis
Successfully installed aioredis-2.0.1 async-timeout-5.0.1

View File

@@ -3,11 +3,10 @@
包含与回调处理相关的业务逻辑函数 包含与回调处理相关的业务逻辑函数
""" """
import json import json
import time import asyncio
import httpx import httpx
from datetime import datetime from datetime import datetime
from typing import Optional, Tuple, Dict, Any from typing import Optional, Tuple, Dict, Any
from sqlalchemy import text
from fastapi import Request from fastapi import Request
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings from app.config import settings
@@ -95,8 +94,8 @@ async def log_callback_request(
return False return False
def save_callback_data_items( async def save_callback_data_items(
conn, db: AsyncSession,
callback_data_items: list, callback_data_items: list,
callback_failure_log_id: int callback_failure_log_id: int
) -> bool: ) -> bool:
@@ -105,12 +104,14 @@ def save_callback_data_items(
Returns: Returns:
bool: 保存是否成功,True表示成功,False表示失败或跳过保存 bool: 保存是否成功,True表示成功,False表示失败或跳过保存
""" """
from sqlalchemy import text
try: try:
logger.info(f"📊 共传入callback_log_id {callback_failure_log_id} 的 {len(callback_data_items)} 条数据") logger.info(f"📊 共传入callback_log_id {callback_failure_log_id} 的 {len(callback_data_items)} 条数据")
# 检查是否已经存在该callback_log_id的数据 # 检查是否已经存在该callback_log_id的数据
existing_data_result = conn.execute( existing_data_result = await db.execute(
text("SELECT * FROM callback_failure_data WHERE callback_failure_log_id = :log_id"), text("SELECT * FROM callback_failure_data WHERE callback_failure_log_id = :log_id"),
{"log_id": callback_failure_log_id} {"log_id": callback_failure_log_id}
) )
@@ -157,7 +158,7 @@ def save_callback_data_items(
raw_data_json = json.dumps(item, ensure_ascii=False) raw_data_json = json.dumps(item, ensure_ascii=False)
# 直接插入数据库 # 直接插入数据库
conn.execute( await db.execute(
text(""" text("""
INSERT INTO callback_failure_data INSERT INTO callback_failure_data
(callback_failure_log_id, phone_number, task_id, user_id, status, status_description, raw_data, calldate) (callback_failure_log_id, phone_number, task_id, user_id, status, status_description, raw_data, calldate)
@@ -175,17 +176,18 @@ def save_callback_data_items(
} }
) )
conn.commit() await db.commit()
logger.info(f"✅ 成功保存 {len(callback_data_items)} 条callback_data记录到数据库") logger.info(f"✅ 成功保存 {len(callback_data_items)} 条callback_data记录到数据库")
return True return True
except Exception as e: except Exception as e:
logger.error(f"❌ 保存callback_data到数据库失败: {e}", exc_info=True) logger.error(f"❌ 保存callback_data到数据库失败: {e}", exc_info=True)
await db.rollback()
# 不重新抛出异常,避免影响主业务流程 # 不重新抛出异常,避免影响主业务流程
return False return False
def get_uncompleted_callback_log(conn) -> Tuple[bool, Optional[Tuple[int, str, str, str]]]: async def get_uncompleted_callback_log(db: AsyncSession) -> Tuple[bool, Optional[Tuple[int, str, str, str]]]:
"""获取一条未完成的回调请求(按创建时间取最小值) """获取一条未完成的回调请求(按创建时间取最小值)
Returns: Returns:
@@ -193,9 +195,11 @@ def get_uncompleted_callback_log(conn) -> Tuple[bool, Optional[Tuple[int, str, s
- 第一个值表示查询是否成功(True表示成功,False表示失败) - 第一个值表示查询是否成功(True表示成功,False表示失败)
- 第二个值为回调日志数据元组或None - 第二个值为回调日志数据元组或None
""" """
from sqlalchemy import text
try: try:
# 查询一条未完成的回调日志(按创建时间升序排列,取第一条) # 查询一条未完成的回调日志(按创建时间升序排列,取第一条)
result = conn.execute( result = await db.execute(
text(""" text("""
SELECT id, site_id, request_headers, request_body SELECT id, site_id, request_headers, request_body
FROM callback_failure_logs FROM callback_failure_logs
@@ -221,11 +225,13 @@ def get_uncompleted_callback_log(conn) -> Tuple[bool, Optional[Tuple[int, str, s
return False, None return False, None
def get_callback_log_data(conn, callback_log_id: int) -> Tuple[bool, Optional[str], Optional[str]]: async def get_callback_log_data(db: AsyncSession, callback_log_id: int) -> Tuple[bool, Optional[str], Optional[str]]:
"""获取回调日志数据""" """获取回调日志数据"""
from sqlalchemy import text
try: try:
# 查询回调日志 # 查询回调日志
result = conn.execute( result = await db.execute(
text("SELECT request_headers, request_body FROM callback_failure_logs WHERE id = :log_id"), text("SELECT request_headers, request_body FROM callback_failure_logs WHERE id = :log_id"),
{"log_id": callback_log_id} {"log_id": callback_log_id}
) )
@@ -245,8 +251,8 @@ def get_callback_log_data(conn, callback_log_id: int) -> Tuple[bool, Optional[st
return False, None, None return False, None, None
def call_external_api_with_retry( async def call_external_api_with_retry(
conn, db: AsyncSession,
request_body: dict, request_body: dict,
request_headers: Dict[str, str], request_headers: Dict[str, str],
max_retries: int, max_retries: int,
@@ -257,16 +263,16 @@ def call_external_api_with_retry(
try: try:
logger.info(f"🌐 尝试调用外部API接口,第{attempt}次") logger.info(f"🌐 尝试调用外部API接口,第{attempt}次")
with httpx.Client(timeout=30.0) as client: async with httpx.AsyncClient(timeout=30.0) as client:
response = client.post( response = await client.post(
settings.external_api_url, settings.external_api_url,
json=request_body, json=request_body,
headers=request_headers headers=request_headers
) )
# 记录推送日志 # 记录推送日志
_log_dtc_push_call( await _log_dtc_push_call(
conn=conn, db=db,
callback_failure_log_id=callback_failure_log_id, callback_failure_log_id=callback_failure_log_id,
request_url=settings.external_api_url, request_url=settings.external_api_url,
request_headers=request_headers, request_headers=request_headers,
@@ -297,18 +303,20 @@ def call_external_api_with_retry(
# 如果不是最后一次尝试,等待一段时间再重试 # 如果不是最后一次尝试,等待一段时间再重试
if attempt < max_retries: if attempt < max_retries:
time.sleep(2 ** attempt) # 指数退避 await asyncio.sleep(2 ** attempt) # 指数退避
# 所有重试都失败了 # 所有重试都失败了
logger.error(f"❌ 调用外部API接口失败,已重试{max_retries}次") logger.error(f"❌ 调用外部API接口失败,已重试{max_retries}次")
return False, max_retries return False, max_retries
def mark_callback_log_completed(conn, callback_log_id: int) -> bool: async def mark_callback_log_completed(db: AsyncSession, callback_log_id: int) -> bool:
"""标记CallbackFailureLog记录为已完成""" """标记CallbackFailureLog记录为已完成"""
from sqlalchemy import text
try: try:
# 先检查推送日志中是否有响应成功的记录 # 先检查推送日志中是否有响应成功的记录
success_result = conn.execute( success_result = await db.execute(
text(""" text("""
SELECT COUNT(*) as success_count SELECT COUNT(*) as success_count
FROM external_api_logs FROM external_api_logs
@@ -326,7 +334,7 @@ def mark_callback_log_completed(conn, callback_log_id: int) -> bool:
logger.info(f"✅ 回调日志 {callback_log_id} 找到 {success_count} 条成功推送记录,开始标记为完成") logger.info(f"✅ 回调日志 {callback_log_id} 找到 {success_count} 条成功推送记录,开始标记为完成")
# 更新指定日志记录为已完成 # 更新指定日志记录为已完成
result = conn.execute( result = await db.execute(
text(""" text("""
UPDATE callback_failure_logs UPDATE callback_failure_logs
SET is_completed = true SET is_completed = true
@@ -341,18 +349,18 @@ def mark_callback_log_completed(conn, callback_log_id: int) -> bool:
logger.warning(f"⚠️ 回调日志 {callback_log_id} 已标记为完成或不存在") logger.warning(f"⚠️ 回调日志 {callback_log_id} 已标记为完成或不存在")
return False return False
conn.commit() await db.commit()
logger.info(f"✅ 回调日志 {callback_log_id} 已标记为完成") logger.info(f"✅ 回调日志 {callback_log_id} 已标记为完成")
return True return True
except Exception as e: except Exception as e:
logger.error(f"❌ 标记回调日志 {callback_log_id} 完成失败: {e}", exc_info=True) logger.error(f"❌ 标记回调日志 {callback_log_id} 完成失败: {e}", exc_info=True)
conn.rollback() await db.rollback()
return False return False
def get_related_records_by_unique_data_list( async def get_related_records_by_unique_data_list(
conn, db: AsyncSession,
unique_data_list: list, unique_data_list: list,
callback_log_id: int, callback_log_id: int,
limit_count: int limit_count: int
@@ -360,7 +368,7 @@ def get_related_records_by_unique_data_list(
"""根据unique_data_list中的数据查询相关记录,按创建时间排序 """根据unique_data_list中的数据查询相关记录,按创建时间排序
Args: Args:
conn: 数据库连接 db: 异步数据库会话
unique_data_list: 包含手机号、task_id、user_id的数据项列表 unique_data_list: 包含手机号、task_id、user_id的数据项列表
callback_log_id: 回调日志ID,作为过滤条件 callback_log_id: 回调日志ID,作为过滤条件
limit_count: 获取记录数量限制 limit_count: 获取记录数量限制
@@ -370,6 +378,8 @@ def get_related_records_by_unique_data_list(
- 第一个值表示查询是否成功(True表示成功,False表示失败) - 第一个值表示查询是否成功(True表示成功,False表示失败)
- 第二个值为查询到的相关记录列表,失败时返回空列表 - 第二个值为查询到的相关记录列表,失败时返回空列表
""" """
from sqlalchemy import text
related_records = [] related_records = []
try: try:
@@ -387,7 +397,7 @@ def get_related_records_by_unique_data_list(
user_id = item.get('user_id', '') user_id = item.get('user_id', '')
if phone_number and task_id and user_id: if phone_number and task_id and user_id:
result = conn.execute( result = await db.execute(
text(""" text("""
SELECT phone_number, task_id, user_id, created_at SELECT phone_number, task_id, user_id, created_at
FROM callback_failure_data FROM callback_failure_data
@@ -427,8 +437,8 @@ def get_related_records_by_unique_data_list(
return False, [] return False, []
def _log_dtc_push_call( async def _log_dtc_push_call(
conn, db: AsyncSession,
callback_failure_log_id: int, callback_failure_log_id: int,
request_url: str, request_url: str,
request_headers: Dict[str, Any], request_headers: Dict[str, Any],
@@ -439,6 +449,8 @@ def _log_dtc_push_call(
retry_count: int retry_count: int
) -> bool: ) -> bool:
"""记录DTC推送调用日志""" """记录DTC推送调用日志"""
from sqlalchemy import text
try: try:
# 提取手机号从请求体中 # 提取手机号从请求体中
phone_number = None phone_number = None
@@ -448,7 +460,7 @@ def _log_dtc_push_call(
phone_number = first_item['number_data']['number'] phone_number = first_item['number_data']['number']
# 记录到external_api_logs表 # 记录到external_api_logs表
conn.execute( await db.execute(
text(""" text("""
INSERT INTO external_api_logs ( INSERT INTO external_api_logs (
callback_failure_log_id, callback_failure_log_id,
@@ -487,11 +499,11 @@ def _log_dtc_push_call(
} }
) )
conn.commit() await db.commit()
logger.debug(f"📝 DTC推送调用日志记录成功,状态码: {response_status}") logger.debug(f"📝 DTC推送调用日志记录成功,状态码: {response_status}")
return True return True
except Exception as e: except Exception as e:
logger.error(f"❌ 记录DTC推送调用日志失败: {e}", exc_info=True) logger.error(f"❌ 记录DTC推送调用日志失败: {e}", exc_info=True)
conn.rollback() await db.rollback()
return False return False

View File

@@ -16,7 +16,7 @@ from app.callback_service import (
mark_callback_log_completed, mark_callback_log_completed,
get_related_records_by_unique_data_list get_related_records_by_unique_data_list
) )
from app.redis_lock import redis_manager from app.redis_lock import redis_manager, async_redis_manager, distributed_lock
import requests import requests
from app.api_config import API_CONFIG, RETRY_CONFIG from app.api_config import API_CONFIG, RETRY_CONFIG
@@ -33,11 +33,18 @@ def conditional_task(enabled=True):
return wrapper return wrapper
return decorator return decorator
# 创建同步数据库连接用于Celery任务 # 创建异步数据库连接用于Celery任务
engine = create_engine( from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
# 使用异步数据库连接
async_engine = create_async_engine(
settings.database_url, settings.database_url,
echo=settings.debug, echo=settings.debug,
future=True future=True,
)
AsyncSessionLocal = async_sessionmaker(
async_engine, class_=AsyncSession, expire_on_commit=False
) )
@celery_app.task(bind=True, name='push_data_to_dtc') @celery_app.task(bind=True, name='push_data_to_dtc')
@@ -47,27 +54,31 @@ def push_data_to_dtc_task(self):
推送数据给DTC的Celery任务 推送数据给DTC的Celery任务
自动获取一条未完成的回调请求进行处理 自动获取一条未完成的回调请求进行处理
""" """
import asyncio
from functools import partial
async def async_task():
task_name = 'push_data_to_dtc' task_name = 'push_data_to_dtc'
logger.info(f"🌿 开始推送数据给DTC任务") logger.info(f"🌿 开始推送数据给DTC任务")
try: try:
# 获取分布式锁,使用任务名称作为锁标识 # 连接异步Redis
connection_success = redis_manager.connect() connection_success = await async_redis_manager.connect()
if not connection_success: if not connection_success:
logger.error(f"❌ Redis管理器连接失败,任务停止执行") logger.error(f"❌ Redis管理器连接失败,任务停止执行")
raise Exception("Redis连接失败") raise Exception("Redis连接失败")
lock = redis_manager.create_lock(f"celery_task:{task_name}", timeout=300) # 5分钟超时 redis_pool = await async_redis_manager.get_redis_pool()
# 尝试获取锁
if not lock.acquire(blocking=False): # 使用异步分布式锁
async with distributed_lock(redis_pool, f"celery_task:{task_name}", timeout=300) as lock_acquired:
if not lock_acquired:
logger.warning(f"⚠️ 任务 {task_name} 正在执行中,跳过本次执行") logger.warning(f"⚠️ 任务 {task_name} 正在执行中,跳过本次执行")
return {"status": "skipped", "message": f"任务 {task_name} 正在执行中,跳过本次执行"} return {"status": "skipped", "message": f"任务 {task_name} 正在执行中,跳过本次执行"}
logger.info(f"🔒 成功获取任务 {task_name} 的分布式锁") logger.info(f"🔒 成功获取任务 {task_name} 的分布式锁")
with engine.connect() as conn: async with AsyncSessionLocal() as db:
logger.info("1")
# 更新任务状态 # 更新任务状态
self.update_state( self.update_state(
state='PROGRESS', state='PROGRESS',
@@ -75,10 +86,10 @@ def push_data_to_dtc_task(self):
) )
logger.info("2") logger.info("2")
# 获取一条未完成的回调请求(按创建时间取最小值) # 获取一条未完成的回调请求(按创建时间取最小值)
query_success, callback_log_data = get_uncompleted_callback_log(conn) query_success, callback_log_data = await get_uncompleted_callback_log(db)
if not query_success: if not query_success:
logger.error("❌ 查询未完成的回调请求失败,尝试再查一次") logger.error("❌ 查询未完成的回调请求失败,尝试再查一次")
query_success, callback_log_data = get_uncompleted_callback_log(conn) query_success, callback_log_data = await get_uncompleted_callback_log(db)
if not query_success: if not query_success:
logger.error("❌ 再次查询未完成的回调请求失败,停止任务执行") logger.error("❌ 再次查询未完成的回调请求失败,停止任务执行")
raise Exception("查询未完成的回调请求失败,任务停止执行") raise Exception("查询未完成的回调请求失败,任务停止执行")
@@ -103,7 +114,7 @@ def push_data_to_dtc_task(self):
# 保存callback_data.data中的数据 # 保存callback_data.data中的数据
data_list = request_body.get('data', []) data_list = request_body.get('data', [])
if data_list and len(data_list) > 0: if data_list and len(data_list) > 0:
save_success = save_callback_data_items(conn, data_list, callback_log_id) save_success = await save_callback_data_items(db, data_list, callback_log_id)
if not save_success: if not save_success:
logger.error(f"❌ 保存callback_data:{callback_log_id}失败,停止任务执行") logger.error(f"❌ 保存callback_data:{callback_log_id}失败,停止任务执行")
raise Exception(f"保存callback_data:{callback_log_id}失败,任务停止执行") raise Exception(f"保存callback_data:{callback_log_id}失败,任务停止执行")
@@ -146,8 +157,8 @@ def push_data_to_dtc_task(self):
logger.info(f"🗑️ 移除了 {original_count - len(unique_data_list)} 条重复数据 - callback_log_id: {callback_log_id}, site_id: {site_id}") logger.info(f"🗑️ 移除了 {original_count - len(unique_data_list)} 条重复数据 - callback_log_id: {callback_log_id}, site_id: {site_id}")
# 查询相关数据:根据手机号、task_id、user_id作为条件,按创建时间排序获取前三条记录 # 查询相关数据:根据手机号、task_id、user_id作为条件,按创建时间排序获取前三条记录
query_success, related_records = get_related_records_by_unique_data_list( query_success, related_records = await get_related_records_by_unique_data_list(
conn, unique_data_list, callback_log_id, settings.count_threshold db, unique_data_list, callback_log_id, settings.count_threshold
) )
if not query_success: if not query_success:
@@ -176,8 +187,8 @@ def push_data_to_dtc_task(self):
meta={'current': 85, 'total': 100, 'status': f'开始推送数据给DTC(ID={callback_log_id}, site_id={site_id})...'} meta={'current': 85, 'total': 100, 'status': f'开始推送数据给DTC(ID={callback_log_id}, site_id={site_id})...'}
) )
success, retry_count = call_external_api_with_retry( success, retry_count = await call_external_api_with_retry(
conn=conn, db=db,
request_body=filtered_request_body, request_body=filtered_request_body,
request_headers=request_headers, request_headers=request_headers,
max_retries=settings.external_api_retry_max, max_retries=settings.external_api_retry_max,
@@ -194,7 +205,7 @@ def push_data_to_dtc_task(self):
logger.info(f"🎉 推送数据给DTC处理完成 - callback_log_id: {callback_log_id}, site_id: {site_id}") logger.info(f"🎉 推送数据给DTC处理完成 - callback_log_id: {callback_log_id}, site_id: {site_id}")
# 标记CallbackFailureLog为已完成 # 标记CallbackFailureLog为已完成
mark_callback_log_completed(conn, callback_log_id) await mark_callback_log_completed(db, callback_log_id)
# 更新任务状态 # 更新任务状态
self.update_state( self.update_state(
@@ -212,13 +223,25 @@ def push_data_to_dtc_task(self):
return {"status": "error", "message": str(e)} return {"status": "error", "message": str(e)}
finally: finally:
# 释放分布式锁 # 分布式锁会通过上下文管理器自动释放
logger.debug(f"🔓 任务 {task_name} 的分布式锁已通过上下文管理器处理")
# 在同步的Celery任务中运行异步代码
try: try:
if 'lock' in locals(): # 创建新的事件循环
lock.release() loop = asyncio.new_event_loop()
logger.info(f"🔓 释放任务 {task_name} 的分布式锁") asyncio.set_event_loop(loop)
return loop.run_until_complete(async_task())
except Exception as e: except Exception as e:
logger.error(f"❌ 释放任务 {task_name} 的分布式锁失败: {e}") logger.error(f"❌ 异步任务执行失败: {e}", exc_info=True)
return {"status": "error", "message": str(e)}
finally:
# 清理事件循环
try:
if 'loop' in locals():
loop.close()
except:
pass
@celery_app.task(bind=True, name='call_api', max_retries=None) @celery_app.task(bind=True, name='call_api', max_retries=None)

View File

@@ -1,15 +1,212 @@
import redis import asyncio
import uuid import uuid
from typing import Optional from typing import Optional
from contextlib import asynccontextmanager
import aioredis
from app.config import settings from app.config import settings
from app.logger import get_redis_logger from app.logger import get_redis_logger
logger = 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: class RedisLock:
"""Redis分布式锁""" """Redis分布式锁 (已废弃,请使用AsyncRedisLock)"""
def __init__(self, redis_client: redis.Redis, key: str, timeout: int = None): def __init__(self, redis_client, key: str, timeout: int = None):
self.redis_client = redis_client self.redis_client = redis_client
self.key = f"lock:{key}" self.key = f"lock:{key}"
self.timeout = timeout or settings.redis_lock_timeout self.timeout = timeout or settings.redis_lock_timeout
@@ -86,14 +283,15 @@ class RedisLock:
class RedisManager: class RedisManager:
def __init__(self): def __init__(self):
self.redis_client: Optional[redis.Redis] = None self.redis_client = None
def connect(self) -> bool: def connect(self) -> bool:
"""连接Redis """连接Redis (已废弃,请使用AsyncRedisManager)
Returns: Returns:
bool: 连接是否成功 bool: 连接是否成功
""" """
import redis
try: try:
logger.info(f"🔴 正在连接Redis: {settings.redis_url}") logger.info(f"🔴 正在连接Redis: {settings.redis_url}")
self.redis_client = redis.from_url( self.redis_client = redis.from_url(
@@ -127,5 +325,6 @@ class RedisManager:
return RedisLock(self.redis_client, key, timeout) return RedisLock(self.redis_client, key, timeout)
# 全局Redis管理器实例 # 全局Redis管理器实例(保持向后兼容)
redis_manager = RedisManager() redis_manager = RedisManager()
async_redis_manager = AsyncRedisManager()

View File

@@ -5,6 +5,7 @@ asyncpg>=0.31.0
alembic>=1.17.2 alembic>=1.17.2
celery>=5.6.0 celery>=5.6.0
redis>=7.1.0 redis>=7.1.0
aioredis==1.3.1
flower>=2.0.1 flower>=2.0.1
pydantic>=2.12.5 pydantic>=2.12.5
pydantic-settings>=2.12.0 pydantic-settings>=2.12.0
@@ -12,5 +13,4 @@ python-multipart>=0.0.20
httpx>=0.28.1 httpx>=0.28.1
python-dotenv>=1.2.1 python-dotenv>=1.2.1
pytest>=9.0.2 pytest>=9.0.2
pytest-asyncio>=1.3.0
requests==2.32.5 requests==2.32.5