From 572cc3113fdd53a30b45cc091de778551fa94126 Mon Sep 17 00:00:00 2001 From: "mark.tian" Date: Wed, 10 Dec 2025 11:40:56 +0800 Subject: [PATCH] =?UTF-8?q?celery=E4=BB=BB=E5=8A=A1=E6=94=B9=E4=B8=BA?= =?UTF-8?q?=E5=BC=82=E6=AD=A5=E8=AE=BF=E9=97=AE=E6=95=B0=E6=8D=AE=E5=BA=93?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- 2.0.0 | 9 + app/callback_service.py | 76 +++++---- app/celery_tasks.py | 369 +++++++++++++++++++++------------------- app/redis_lock.py | 213 ++++++++++++++++++++++- requirements.txt | 2 +- 5 files changed, 456 insertions(+), 213 deletions(-) create mode 100644 2.0.0 diff --git a/2.0.0 b/2.0.0 new file mode 100644 index 0000000..529590c --- /dev/null +++ b/2.0.0 @@ -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 diff --git a/app/callback_service.py b/app/callback_service.py index cf911ac..c0b2474 100644 --- a/app/callback_service.py +++ b/app/callback_service.py @@ -3,11 +3,10 @@ 包含与回调处理相关的业务逻辑函数 """ import json -import time +import asyncio import httpx from datetime import datetime from typing import Optional, Tuple, Dict, Any -from sqlalchemy import text from fastapi import Request from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings @@ -95,8 +94,8 @@ async def log_callback_request( return False -def save_callback_data_items( - conn, +async def save_callback_data_items( + db: AsyncSession, callback_data_items: list, callback_failure_log_id: int ) -> bool: @@ -105,12 +104,14 @@ def save_callback_data_items( Returns: bool: 保存是否成功,True表示成功,False表示失败或跳过保存 """ + from sqlalchemy import text + try: logger.info(f"📊 共传入callback_log_id {callback_failure_log_id} 的 {len(callback_data_items)} 条数据") # 检查是否已经存在该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"), {"log_id": callback_failure_log_id} ) @@ -157,7 +158,7 @@ def save_callback_data_items( raw_data_json = json.dumps(item, ensure_ascii=False) # 直接插入数据库 - conn.execute( + await db.execute( text(""" INSERT INTO callback_failure_data (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记录到数据库") return True except Exception as e: logger.error(f"❌ 保存callback_data到数据库失败: {e}", exc_info=True) + await db.rollback() # 不重新抛出异常,避免影响主业务流程 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: @@ -193,9 +195,11 @@ def get_uncompleted_callback_log(conn) -> Tuple[bool, Optional[Tuple[int, str, s - 第一个值表示查询是否成功(True表示成功,False表示失败) - 第二个值为回调日志数据元组或None """ + from sqlalchemy import text + try: # 查询一条未完成的回调日志(按创建时间升序排列,取第一条) - result = conn.execute( + result = await db.execute( text(""" SELECT id, site_id, request_headers, request_body FROM callback_failure_logs @@ -221,11 +225,13 @@ def get_uncompleted_callback_log(conn) -> Tuple[bool, Optional[Tuple[int, str, s 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: # 查询回调日志 - result = conn.execute( + result = await db.execute( text("SELECT request_headers, request_body FROM callback_failure_logs WHERE id = :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 -def call_external_api_with_retry( - conn, +async def call_external_api_with_retry( + db: AsyncSession, request_body: dict, request_headers: Dict[str, str], max_retries: int, @@ -257,16 +263,16 @@ def call_external_api_with_retry( try: logger.info(f"🌐 尝试调用外部API接口,第{attempt}次") - with httpx.Client(timeout=30.0) as client: - response = client.post( + async with httpx.AsyncClient(timeout=30.0) as client: + response = await client.post( settings.external_api_url, json=request_body, headers=request_headers ) # 记录推送日志 - _log_dtc_push_call( - conn=conn, + await _log_dtc_push_call( + db=db, callback_failure_log_id=callback_failure_log_id, request_url=settings.external_api_url, request_headers=request_headers, @@ -297,18 +303,20 @@ def call_external_api_with_retry( # 如果不是最后一次尝试,等待一段时间再重试 if attempt < max_retries: - time.sleep(2 ** attempt) # 指数退避 + await asyncio.sleep(2 ** attempt) # 指数退避 # 所有重试都失败了 logger.error(f"❌ 调用外部API接口失败,已重试{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记录为已完成""" + from sqlalchemy import text + try: # 先检查推送日志中是否有响应成功的记录 - success_result = conn.execute( + success_result = await db.execute( text(""" SELECT COUNT(*) as success_count 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} 条成功推送记录,开始标记为完成") # 更新指定日志记录为已完成 - result = conn.execute( + result = await db.execute( text(""" UPDATE callback_failure_logs 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} 已标记为完成或不存在") return False - conn.commit() + await db.commit() logger.info(f"✅ 回调日志 {callback_log_id} 已标记为完成") return True except Exception as e: logger.error(f"❌ 标记回调日志 {callback_log_id} 完成失败: {e}", exc_info=True) - conn.rollback() + await db.rollback() return False -def get_related_records_by_unique_data_list( - conn, +async def get_related_records_by_unique_data_list( + db: AsyncSession, unique_data_list: list, callback_log_id: int, limit_count: int @@ -360,7 +368,7 @@ def get_related_records_by_unique_data_list( """根据unique_data_list中的数据查询相关记录,按创建时间排序 Args: - conn: 数据库连接 + db: 异步数据库会话 unique_data_list: 包含手机号、task_id、user_id的数据项列表 callback_log_id: 回调日志ID,作为过滤条件 limit_count: 获取记录数量限制 @@ -370,6 +378,8 @@ def get_related_records_by_unique_data_list( - 第一个值表示查询是否成功(True表示成功,False表示失败) - 第二个值为查询到的相关记录列表,失败时返回空列表 """ + from sqlalchemy import text + related_records = [] try: @@ -387,7 +397,7 @@ def get_related_records_by_unique_data_list( user_id = item.get('user_id', '') if phone_number and task_id and user_id: - result = conn.execute( + result = await db.execute( text(""" SELECT phone_number, task_id, user_id, created_at FROM callback_failure_data @@ -427,8 +437,8 @@ def get_related_records_by_unique_data_list( return False, [] -def _log_dtc_push_call( - conn, +async def _log_dtc_push_call( + db: AsyncSession, callback_failure_log_id: int, request_url: str, request_headers: Dict[str, Any], @@ -439,6 +449,8 @@ def _log_dtc_push_call( retry_count: int ) -> bool: """记录DTC推送调用日志""" + from sqlalchemy import text + try: # 提取手机号从请求体中 phone_number = None @@ -448,7 +460,7 @@ def _log_dtc_push_call( phone_number = first_item['number_data']['number'] # 记录到external_api_logs表 - conn.execute( + await db.execute( text(""" INSERT INTO external_api_logs ( callback_failure_log_id, @@ -487,11 +499,11 @@ def _log_dtc_push_call( } ) - conn.commit() + await db.commit() logger.debug(f"📝 DTC推送调用日志记录成功,状态码: {response_status}") return True except Exception as e: logger.error(f"❌ 记录DTC推送调用日志失败: {e}", exc_info=True) - conn.rollback() + await db.rollback() return False \ No newline at end of file diff --git a/app/celery_tasks.py b/app/celery_tasks.py index 4a09131..6bf026e 100644 --- a/app/celery_tasks.py +++ b/app/celery_tasks.py @@ -16,7 +16,7 @@ from app.callback_service import ( mark_callback_log_completed, 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 from app.api_config import API_CONFIG, RETRY_CONFIG @@ -33,11 +33,18 @@ def conditional_task(enabled=True): return wrapper return decorator -# 创建同步数据库连接用于Celery任务 -engine = create_engine( +# 创建异步数据库连接用于Celery任务 +from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker + +# 使用异步数据库连接 +async_engine = create_async_engine( settings.database_url, 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') @@ -47,178 +54,194 @@ def push_data_to_dtc_task(self): 推送数据给DTC的Celery任务 自动获取一条未完成的回调请求进行处理 """ - task_name = 'push_data_to_dtc' - logger.info(f"🌿 开始推送数据给DTC任务") - - try: - # 获取分布式锁,使用任务名称作为锁标识 - connection_success = redis_manager.connect() - if not connection_success: - logger.error(f"❌ Redis管理器连接失败,任务停止执行") - raise Exception("Redis连接失败") + import asyncio + from functools import partial + + async def async_task(): + task_name = 'push_data_to_dtc' + logger.info(f"🌿 开始推送数据给DTC任务") - lock = redis_manager.create_lock(f"celery_task:{task_name}", timeout=300) # 5分钟超时 - # 尝试获取锁 - if not lock.acquire(blocking=False): - logger.warning(f"⚠️ 任务 {task_name} 正在执行中,跳过本次执行") - return {"status": "skipped", "message": f"任务 {task_name} 正在执行中,跳过本次执行"} - - logger.info(f"🔒 成功获取任务 {task_name} 的分布式锁") - - with engine.connect() as conn: - - logger.info("1") - # 更新任务状态 - self.update_state( - state='PROGRESS', - meta={'current': 0, 'total': 100, 'status': f'正在获取下一条回调请求日志...'} - ) - logger.info("2") - # 获取一条未完成的回调请求(按创建时间取最小值) - query_success, callback_log_data = get_uncompleted_callback_log(conn) - if not query_success: - logger.error("❌ 查询未完成的回调请求失败,尝试再查一次") - query_success, callback_log_data = get_uncompleted_callback_log(conn) - if not query_success: - logger.error("❌ 再次查询未完成的回调请求失败,停止任务执行") - raise Exception("查询未完成的回调请求失败,任务停止执行") - logger.info("3") - if not callback_log_data: - logger.info("📋 没有找到未完成的回调请求") - return {"status": "skipped", "message": "没有找到未完成的回调请求"} - - callback_log_id, site_id, request_headers_json, request_body_json = callback_log_data - logger.info(f"📋 获取到未完成的回调请求: ID={callback_log_id}, site_id={site_id}") - - # 更新任务状态 - self.update_state( - state='PROGRESS', - meta={'current': 25, 'total': 100, 'status': f'获取回调请求成功(ID={callback_log_id}, site_id={site_id}),开始分析处理...'} - ) - - # 解析请求头和请求体 - request_headers = json.loads(request_headers_json) - request_body = json.loads(request_body_json) - - # 保存callback_data.data中的数据 - data_list = request_body.get('data', []) - if data_list and len(data_list) > 0: - save_success = save_callback_data_items(conn, data_list, callback_log_id) - if not save_success: - logger.error(f"❌ 保存callback_data:{callback_log_id}失败,停止任务执行") - raise Exception(f"保存callback_data:{callback_log_id}失败,任务停止执行") - - # 更新任务状态 - self.update_state( - state='PROGRESS', - meta={'current': 50, 'total': 100, 'status': f'请求体中通话明细处理成功(ID={callback_log_id}, site_id={site_id}),开始过滤需要转发的通话记录...'} - ) - - # 根据手机号、task_id、user_id给data_list去重 - data_list = request_body.get('data', []) - unique_data_list = [] - seen_records = set() - original_count = len(data_list) - - for item in data_list: - # 提取手机号 - number_data = item.get('number_data', {}) - phone_number = number_data.get('number', '') - - # 提取task_id - task = item.get('task', {}) - task_id = task.get('id', '') - - # 提取user_id - user_id = item.get('user_id', '') - - # 创建唯一标识 - unique_key = (phone_number, task_id, user_id) - - # 如果这个组合没见过,则添加到去重列表中 - if unique_key not in seen_records: - seen_records.add(unique_key) - unique_data_list.append(item) - - logger.info(f"🔄 数据去重完成: 原始数据 {original_count} 条,去重后 {len(unique_data_list)} 条 - callback_log_id: {callback_log_id}, site_id: {site_id}") - - if len(unique_data_list) < original_count: - logger.info(f"🗑️ 移除了 {original_count - len(unique_data_list)} 条重复数据 - callback_log_id: {callback_log_id}, site_id: {site_id}") - - # 查询相关数据:根据手机号、task_id、user_id作为条件,按创建时间排序获取前三条记录 - query_success, related_records = get_related_records_by_unique_data_list( - conn, unique_data_list, callback_log_id, settings.count_threshold - ) - - if not query_success: - logger.error("❌ 查询相关记录失败,停止任务执行") - raise Exception("查询相关记录失败,任务停止执行") - - # 判断是否有相关记录需要处理 - if not related_records: - logger.info(f"📋 没有有效的通过记录需要处理,直接返回 - callback_log_id: {callback_log_id}, site_id={site_id}") - return {"status": "completed", "message": "没有有效的通过记录需要处理"} - - # 更新任务状态 - self.update_state( - state='PROGRESS', - meta={'current': 70, 'total': 100, 'status': f'需要推送的通过记录已获取成功(ID={callback_log_id}, site_id={site_id}),准备转发...'} - ) - - if related_records: - # 创建只包含有效数据项的请求体 - filtered_request_body = request_body.copy() - filtered_request_body['data'] = related_records - - # 更新任务状态 - self.update_state( - state='PROGRESS', - meta={'current': 85, 'total': 100, 'status': f'开始推送数据给DTC(ID={callback_log_id}, site_id={site_id})...'} - ) - - success, retry_count = call_external_api_with_retry( - conn=conn, - request_body=filtered_request_body, - request_headers=request_headers, - max_retries=settings.external_api_retry_max, - callback_failure_log_id=callback_log_id - ) - - if success: - logger.info(f"✅ 推送数据给DTC成功,处理的数据项数量: {len(related_records)}, 重试次数: {retry_count} - callback_log_id: {callback_log_id}, site_id: {site_id}") - else: - logger.error(f"❌ 推送数据给DTC失败,重试次数: {retry_count} - callback_log_id: {callback_log_id}, site_id: {site_id}") - else: - logger.info(f"📋 没有有效的通过记录需要处理 - 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为已完成 - mark_callback_log_completed(conn, callback_log_id) - - # 更新任务状态 - self.update_state( - state='PROGRESS', - meta={'current': 100, 'total': 100, 'status': f'标记回调请求日志为已完成(ID={callback_log_id}, site_id={site_id})'} - ) - - return { - "status": "completed", - "message": "任务完成" - } - - except Exception as e: - logger.error(f"❌ 任务执行失败: {e}", exc_info=True) - return {"status": "error", "message": str(e)} - - finally: - # 释放分布式锁 try: - if 'lock' in locals(): - lock.release() - logger.info(f"🔓 释放任务 {task_name} 的分布式锁") + # 连接异步Redis + connection_success = await async_redis_manager.connect() + if not connection_success: + logger.error(f"❌ Redis管理器连接失败,任务停止执行") + raise Exception("Redis连接失败") + + redis_pool = await async_redis_manager.get_redis_pool() + + # 使用异步分布式锁 + 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} 正在执行中,跳过本次执行") + return {"status": "skipped", "message": f"任务 {task_name} 正在执行中,跳过本次执行"} + + logger.info(f"🔒 成功获取任务 {task_name} 的分布式锁") + + async with AsyncSessionLocal() as db: + # 更新任务状态 + self.update_state( + state='PROGRESS', + meta={'current': 0, 'total': 100, 'status': f'正在获取下一条回调请求日志...'} + ) + logger.info("2") + # 获取一条未完成的回调请求(按创建时间取最小值) + query_success, callback_log_data = await get_uncompleted_callback_log(db) + if not query_success: + logger.error("❌ 查询未完成的回调请求失败,尝试再查一次") + query_success, callback_log_data = await get_uncompleted_callback_log(db) + if not query_success: + logger.error("❌ 再次查询未完成的回调请求失败,停止任务执行") + raise Exception("查询未完成的回调请求失败,任务停止执行") + logger.info("3") + if not callback_log_data: + logger.info("📋 没有找到未完成的回调请求") + return {"status": "skipped", "message": "没有找到未完成的回调请求"} + + callback_log_id, site_id, request_headers_json, request_body_json = callback_log_data + logger.info(f"📋 获取到未完成的回调请求: ID={callback_log_id}, site_id={site_id}") + + # 更新任务状态 + self.update_state( + state='PROGRESS', + meta={'current': 25, 'total': 100, 'status': f'获取回调请求成功(ID={callback_log_id}, site_id={site_id}),开始分析处理...'} + ) + + # 解析请求头和请求体 + request_headers = json.loads(request_headers_json) + request_body = json.loads(request_body_json) + + # 保存callback_data.data中的数据 + data_list = request_body.get('data', []) + if data_list and len(data_list) > 0: + save_success = await save_callback_data_items(db, data_list, callback_log_id) + if not save_success: + logger.error(f"❌ 保存callback_data:{callback_log_id}失败,停止任务执行") + raise Exception(f"保存callback_data:{callback_log_id}失败,任务停止执行") + + # 更新任务状态 + self.update_state( + state='PROGRESS', + meta={'current': 50, 'total': 100, 'status': f'请求体中通话明细处理成功(ID={callback_log_id}, site_id={site_id}),开始过滤需要转发的通话记录...'} + ) + + # 根据手机号、task_id、user_id给data_list去重 + data_list = request_body.get('data', []) + unique_data_list = [] + seen_records = set() + original_count = len(data_list) + + for item in data_list: + # 提取手机号 + number_data = item.get('number_data', {}) + phone_number = number_data.get('number', '') + + # 提取task_id + task = item.get('task', {}) + task_id = task.get('id', '') + + # 提取user_id + user_id = item.get('user_id', '') + + # 创建唯一标识 + unique_key = (phone_number, task_id, user_id) + + # 如果这个组合没见过,则添加到去重列表中 + if unique_key not in seen_records: + seen_records.add(unique_key) + unique_data_list.append(item) + + logger.info(f"🔄 数据去重完成: 原始数据 {original_count} 条,去重后 {len(unique_data_list)} 条 - callback_log_id: {callback_log_id}, site_id: {site_id}") + + if len(unique_data_list) < original_count: + logger.info(f"🗑️ 移除了 {original_count - len(unique_data_list)} 条重复数据 - callback_log_id: {callback_log_id}, site_id: {site_id}") + + # 查询相关数据:根据手机号、task_id、user_id作为条件,按创建时间排序获取前三条记录 + query_success, related_records = await get_related_records_by_unique_data_list( + db, unique_data_list, callback_log_id, settings.count_threshold + ) + + if not query_success: + logger.error("❌ 查询相关记录失败,停止任务执行") + raise Exception("查询相关记录失败,任务停止执行") + + # 判断是否有相关记录需要处理 + if not related_records: + logger.info(f"📋 没有有效的通过记录需要处理,直接返回 - callback_log_id: {callback_log_id}, site_id={site_id}") + return {"status": "completed", "message": "没有有效的通过记录需要处理"} + + # 更新任务状态 + self.update_state( + state='PROGRESS', + meta={'current': 70, 'total': 100, 'status': f'需要推送的通过记录已获取成功(ID={callback_log_id}, site_id={site_id}),准备转发...'} + ) + + if related_records: + # 创建只包含有效数据项的请求体 + filtered_request_body = request_body.copy() + filtered_request_body['data'] = related_records + + # 更新任务状态 + self.update_state( + state='PROGRESS', + meta={'current': 85, 'total': 100, 'status': f'开始推送数据给DTC(ID={callback_log_id}, site_id={site_id})...'} + ) + + success, retry_count = await call_external_api_with_retry( + db=db, + request_body=filtered_request_body, + request_headers=request_headers, + max_retries=settings.external_api_retry_max, + callback_failure_log_id=callback_log_id + ) + + if success: + logger.info(f"✅ 推送数据给DTC成功,处理的数据项数量: {len(related_records)}, 重试次数: {retry_count} - callback_log_id: {callback_log_id}, site_id: {site_id}") + else: + logger.error(f"❌ 推送数据给DTC失败,重试次数: {retry_count} - callback_log_id: {callback_log_id}, site_id: {site_id}") + else: + logger.info(f"📋 没有有效的通过记录需要处理 - 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为已完成 + await mark_callback_log_completed(db, callback_log_id) + + # 更新任务状态 + self.update_state( + state='PROGRESS', + meta={'current': 100, 'total': 100, 'status': f'标记回调请求日志为已完成(ID={callback_log_id}, site_id={site_id})'} + ) + + return { + "status": "completed", + "message": "任务完成" + } + except Exception as e: - logger.error(f"❌ 释放任务 {task_name} 的分布式锁失败: {e}") + logger.error(f"❌ 任务执行失败: {e}", exc_info=True) + return {"status": "error", "message": str(e)} + + finally: + # 分布式锁会通过上下文管理器自动释放 + logger.debug(f"🔓 任务 {task_name} 的分布式锁已通过上下文管理器处理") + + # 在同步的Celery任务中运行异步代码 + try: + # 创建新的事件循环 + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + return loop.run_until_complete(async_task()) + except Exception as 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) diff --git a/app/redis_lock.py b/app/redis_lock.py index 7bff8e3..f024850 100644 --- a/app/redis_lock.py +++ b/app/redis_lock.py @@ -1,15 +1,212 @@ -import redis +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分布式锁""" - def __init__(self, redis_client: redis.Redis, key: str, timeout: int = None): + """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 @@ -86,14 +283,15 @@ class RedisLock: class RedisManager: def __init__(self): - self.redis_client: Optional[redis.Redis] = None + self.redis_client = None def connect(self) -> bool: - """连接Redis + """连接Redis (已废弃,请使用AsyncRedisManager) Returns: bool: 连接是否成功 """ + import redis try: logger.info(f"🔴 正在连接Redis: {settings.redis_url}") self.redis_client = redis.from_url( @@ -127,5 +325,6 @@ class RedisManager: return RedisLock(self.redis_client, key, timeout) -# 全局Redis管理器实例 -redis_manager = RedisManager() \ No newline at end of file +# 全局Redis管理器实例(保持向后兼容) +redis_manager = RedisManager() +async_redis_manager = AsyncRedisManager() \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 2d60bb8..75794c6 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,6 +5,7 @@ asyncpg>=0.31.0 alembic>=1.17.2 celery>=5.6.0 redis>=7.1.0 +aioredis==1.3.1 flower>=2.0.1 pydantic>=2.12.5 pydantic-settings>=2.12.0 @@ -12,5 +13,4 @@ python-multipart>=0.0.20 httpx>=0.28.1 python-dotenv>=1.2.1 pytest>=9.0.2 -pytest-asyncio>=1.3.0 requests==2.32.5 \ No newline at end of file