diff --git a/.env.example b/.env.example index 4589b99..c1572a3 100644 --- a/.env.example +++ b/.env.example @@ -5,6 +5,16 @@ DATABASE_URL=postgresql+asyncpg://user:password@localhost:5432/ai_talk_callback_ CELERY_BROKER_URL=redis://localhost:6379/0 CELERY_RESULT_BACKEND=redis://localhost:6379/0 +# Redis配置 (扩展配置,如果需要覆盖默认值) +REDIS_PASSWORD= +REDIS_MAX_CONNECTIONS=20 +REDIS_TIMEOUT=5 + +# Redis分布式锁配置 +REDIS_LOCK_TIMEOUT=300 +REDIS_LOCK_MAX_RETRIES=10 +REDIS_LOCK_RETRY_DELAY=0.5 + # 业务配置 COUNT_THRESHOLD=3 EXTERNAL_API_ENABLED=false diff --git a/app/celery_tasks.py b/app/celery_tasks.py index 6f39ea4..66dca78 100644 --- a/app/celery_tasks.py +++ b/app/celery_tasks.py @@ -2,409 +2,174 @@ Celery任务定义 """ from celery import current_task -from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker -from sqlalchemy import text +from sqlalchemy import create_engine, text import json -import httpx -import asyncio from app.celery_app import celery_app from app.config import settings from app.logger import get_logger from app.database import CallbackFailureLog, ExternalApiLog, CallbackFailureData +from app.callback_service import ( + save_callback_data_items, + get_uncompleted_callback_log, + get_callback_log_data, + check_phone_number_threshold, + call_external_api_with_retry, + mark_callback_log_completed, + _log_dtc_push_call +) +from app.redis_lock import redis_manager logger = get_logger("celery_tasks") -# 创建独立的数据库连接用于Celery任务 -engine = create_async_engine( +# 创建同步数据库连接用于Celery任务 +engine = create_engine( settings.database_url, echo=settings.debug, future=True ) -AsyncSessionLocal = async_sessionmaker( - engine, - class_=AsyncSession, - expire_on_commit=False -) - - -def get_db(): - """获取数据库会话""" - return AsyncSessionLocal() - - @celery_app.task(bind=True, name='push_data_to_dtc') def push_data_to_dtc_task(self): """ 推送数据给DTC的Celery任务 自动获取一条未完成的回调请求进行处理 """ + task_name = 'push_data_to_dtc' logger.info(f"🌿 开始推送数据给DTC任务") - # 使用asyncio运行异步逻辑 - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) + # 获取分布式锁,使用任务名称作为锁标识 + lock = redis_manager.create_lock(f"celery_task:{task_name}", timeout=300) # 5分钟超时 try: - result = loop.run_until_complete( - _push_data_to_dtc_async(self.request.id) - ) - return result - except Exception as e: - logger.error(f"❌ 推送数据给DTC任务执行失败: {e}", exc_info=True) - raise - finally: - loop.close() - - -async def _push_data_to_dtc_async(task_id: str): - """异步推送数据给DTC的核心逻辑""" - - async with AsyncSessionLocal() as db: + # 尝试获取锁 + if not lock.acquire(blocking=False): + logger.warning(f"⚠️ 任务 {task_name} 正在执行中,跳过本次执行") + return {"status": "skipped", "message": f"任务 {task_name} 正在执行中,跳过本次执行"} + + logger.info(f"🔒 成功获取任务 {task_name} 的分布式锁") + try: - # 获取一条未完成的回调请求(按创建时间取最小值) - callback_log_data = await get_uncompleted_callback_log(db) - 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}") - - # 解析请求头和请求体 - 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: - await save_callback_data_items(db, data_list, callback_log_id) - - # 提取所有手机号并去重 - phone_numbers_set = set() - data_list = request_body.get('data', []) - for item in data_list: - number_data = item.get('number_data', {}) - phone_number = number_data.get('number') - if phone_number: - phone_numbers_set.add(phone_number) - - phone_numbers = list(phone_numbers_set) - logger.info(f"📱 提取到的手机号列表(去重后): {phone_numbers}") - - # 检查手机号列表是否为空 - if not phone_numbers or len(phone_numbers) == 0: - logger.warning(f"⚠️ 手机号列表为空,跳过推送数据给DTC") - return {"status": "skipped", "message": "手机号列表为空"} - - # 处理每个手机号 - processed_count = 0 - skipped_count = 0 - - for phone_number in phone_numbers: + with engine.connect() as conn: + # 获取一条未完成的回调请求(按创建时间取最小值) + 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("查询未完成的回调请求失败,任务停止执行") + + 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}") # 更新任务状态 current_task.update_state( state='PROGRESS', - meta={'current': processed_count + skipped_count, 'total': len(phone_numbers), 'status': f'处理手机号: {phone_number}'} + meta={'current': len(valid_phone_numbers), 'total': len(phone_numbers), 'status': f'检查手机号: {phone_number}'} ) - # 检查手机号是否超过阈值 - exceeds_threshold, query_success = await check_phone_number_threshold(db, phone_number) - - if not query_success: - logger.error(f"❌ 查询手机号 {phone_number} 失败,跳过处理") - skipped_count += 1 - continue - - if exceeds_threshold: - logger.info(f"✅ 手机号 {phone_number} 出现次数超过阈值,跳过处理") - skipped_count += 1 - continue + # 解析请求头和请求体 + request_headers = json.loads(request_headers_json) + request_body = json.loads(request_body_json) - # 推送数据给DTC - success, retry_count = await push_data_to_dtc_with_retry( - db=db, - request_body=request_body, - max_retries=settings.external_api_retry_max, - callback_failure_log_id=callback_log_id - ) + # 保存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}失败,任务停止执行") + + # 提取所有手机号并去重 + phone_numbers_set = set() + data_list = request_body.get('data', []) + for item in data_list: + number_data = item.get('number_data', {}) + phone_number = number_data.get('number') + if phone_number: + phone_numbers_set.add(phone_number) - if success: - logger.info(f"✅ 推送数据给DTC成功,手机号: {phone_number}, 重试次数: {retry_count}") - processed_count += 1 + phone_numbers = list(phone_numbers_set) + logger.info(f"📱 提取到的手机号列表(去重后): {phone_numbers}") + + # 过滤出需要处理的手机号(不超过阈值的) + valid_phone_numbers = [] + for phone_number in phone_numbers: + # 更新任务状态 + current_task.update_state( + state='PROGRESS', + meta={'current': len(valid_phone_numbers), 'total': len(phone_numbers), 'status': f'检查手机号: {phone_number}'} + ) + + # 检查手机号是否超过阈值 + exceeds_threshold, query_success = check_phone_number_threshold(conn, phone_number) + + if not query_success: + logger.error(f"❌ 查询手机号 {phone_number} 失败,跳过处理") + continue + + if exceeds_threshold: + logger.info(f"✅ 手机号 {phone_number} 出现次数超过阈值,跳过处理") + continue + + valid_phone_numbers.append(phone_number) + + logger.info(f"📱 需要处理的手机号数量: {len(valid_phone_numbers)}") + + # 推送数据给DTC(在循环外执行一次) + processed_count = 0 + skipped_count = len(phone_numbers) - len(valid_phone_numbers) + + if valid_phone_numbers: + # 更新任务状态 + current_task.update_state( + state='PROGRESS', + meta={'current': 0, 'total': len(valid_phone_numbers), 'status': '开始推送数据给DTC'} + ) + + success, retry_count = call_external_api_with_retry( + conn=conn, + request_body=request_body, + max_retries=settings.external_api_retry_max, + callback_failure_log_id=callback_log_id + ) + + if success: + logger.info(f"✅ 推送数据给DTC成功,处理的手机号数量: {len(valid_phone_numbers)}, 重试次数: {retry_count}") + processed_count = len(valid_phone_numbers) + else: + logger.error(f"❌ 推送数据给DTC失败,重试次数: {retry_count}") else: - logger.error(f"❌ 推送数据给DTC失败,手机号: {phone_number}, 重试次数: {retry_count}") - skipped_count += 1 + logger.info("📋 没有有效的手机号需要处理") - logger.info(f"🎉 推送数据给DTC处理完成,成功: {processed_count}, 跳过: {skipped_count}") - - # 标记CallbackFailureLog为已完成 - await mark_callback_log_completed(db, callback_log_id) - - return { - "status": "completed", - "processed": processed_count, - "skipped": skipped_count, - "total": len(phone_numbers) - } - + logger.info(f"🎉 推送数据给DTC处理完成,成功: {processed_count}, 跳过: {skipped_count}") + + # 标记CallbackFailureLog为已完成 + mark_callback_log_completed(conn, callback_log_id) + + return { + "status": "completed", + "processed": processed_count, + "skipped": skipped_count, + "total": len(phone_numbers) + } + except Exception as e: logger.error(f"❌ 推送数据给DTC时发生错误: {e}", exc_info=True) return {"status": "error", "message": str(e)} - - -async def save_callback_data_items( - db: AsyncSession, - callback_data_items: list, - callback_failure_log_id: int -): - """保存callback_data.data中的数据到数据库""" - from datetime import datetime + + finally: + # 释放分布式锁 + try: + lock.release() + logger.info(f"🔓 释放任务 {task_name} 的分布式锁") + except Exception as e: + logger.error(f"❌ 释放任务 {task_name} 的分布式锁失败: {e}") - try: - callback_data_records = [] - for item in callback_data_items: - # 获取手机号 - number_data = item.get('number_data', {}) - phone_number = number_data.get('number') - - if not phone_number: - continue - - # 获取任务ID - task = item.get('task', {}) - task_id = task.get('id', '') - - # 获取用户ID - user_id = item.get('user_id', '') - - # 获取状态信息 - status = item.get('status', 0) - status_description = item.get('status_str', '') - - # 获取通话日期,默认使用当前时间 - calldate = datetime.now() - if 'calldate' in item: - try: - # 如果calldate是字符串,尝试解析为datetime - if isinstance(item['calldate'], str): - calldate = datetime.fromisoformat(item['calldate'].replace('Z', '+00:00')) - elif isinstance(item['calldate'], (int, float)): - # 如果是时间戳,转换为datetime - calldate = datetime.fromtimestamp(item['calldate']) - except (ValueError, TypeError) as e: - logger.warning(f"⚠️ 解析calldate失败: {item.get('calldate')}, 使用当前时间, 错误: {e}") - calldate = datetime.now() - - # 将整个item转换为JSON字符串保存 - raw_data_json = json.dumps(item, ensure_ascii=False) - callback_data_record = CallbackFailureData( - callback_failure_log_id=callback_failure_log_id, - phone_number=phone_number, - task_id=task_id, - user_id=user_id, # 保存用户ID - status=status, - status_description=status_description, - raw_data=raw_data_json, # 保存原始JSON字符串 - calldate=calldate # 保存通话日期 - ) - callback_data_records.append(callback_data_record) - - if len(callback_data_records) > 0: - db.add_all(callback_data_records) - await db.commit() - - logger.info(f"✅ 成功保存 {len(callback_data_records)} 条callback_data记录到数据库") - except Exception as e: - logger.error(f"❌ 保存callback_data到数据库失败: {e}", exc_info=True) - # 不重新抛出异常,避免影响主业务流程 + logger.error(f"❌ 获取分布式锁失败: {e}", exc_info=True) + return {"status": "error", "message": f"获取分布式锁失败: {str(e)}"} - -async def get_uncompleted_callback_log(db: AsyncSession): - """获取一条未完成的回调请求(按创建时间取最小值)""" - try: - # 查询一条未完成的回调日志(按创建时间升序排列,取第一条) - result = await db.execute( - text(""" - SELECT id, site_id, request_headers, request_body - FROM callback_failure_logs - WHERE is_completed = false - ORDER BY created_at ASC - LIMIT 1 - """) - ) - row = result.fetchone() - - if not row: - logger.info("📋 没有找到未完成的回调日志") - return None - - callback_log_id, site_id, request_headers, request_body = row - request_headers_json = json.dumps(request_headers, ensure_ascii=False) - request_body_json = json.dumps(request_body, ensure_ascii=False) - - return callback_log_id, site_id, request_headers_json, request_body_json - - except Exception as e: - logger.error(f"❌ 查询未完成回调日志失败: {e}", exc_info=True) - return None - - -async def get_callback_log_data(db: AsyncSession, callback_log_id: int): - """获取回调日志数据""" - try: - # 查询回调日志 - result = await db.execute( - text("SELECT request_headers, request_body FROM callback_failure_logs WHERE id = :log_id"), - {"log_id": callback_log_id} - ) - row = result.fetchone() - - if not row: - logger.error(f"❌ 未找到回调日志: {callback_log_id}") - return False, None, None - - request_headers_json = json.dumps(row[0], ensure_ascii=False) - request_body_json = json.dumps(row[1], ensure_ascii=False) - - return True, request_headers_json, request_body_json - - except Exception as e: - logger.error(f"❌ 查询回调日志失败: {e}", exc_info=True) - return False, None, None - - -async def check_phone_number_threshold(db: AsyncSession, phone_number: str): - """检查手机号出现次数是否超过阈值""" - try: - # 查询手机号在callback_failure_data表中的出现次数 - result = await db.execute( - text("SELECT COUNT(*) FROM callback_failure_data WHERE phone_number = :phone_number"), - {"phone_number": phone_number} - ) - count = result.scalar() - - exceeds_threshold = count >= settings.count_threshold - logger.info(f"📊 手机号 {phone_number} 出现次数: {count}, 阈值: {settings.count_threshold}, 超过阈值: {exceeds_threshold}") - - return exceeds_threshold, True - - except Exception as e: - logger.error(f"❌ 查询手机号 {phone_number} 失败: {e}", exc_info=True) - return False, False - - -async def push_data_to_dtc_with_retry( - db: AsyncSession, - request_body: dict, - max_retries: int, - callback_failure_log_id: int -): - """推送数据给DTC并支持重试""" - - for attempt in range(1, max_retries + 1): - try: - logger.info(f"🌐 尝试推送数据给DTC,第{attempt}次") - - async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.post( - settings.external_api_url, - json=request_body, - headers={"Content-Type": "application/json"} - ) - - # 记录推送日志 - await _log_dtc_push_call( - db=db, - callback_failure_log_id=callback_failure_log_id, - request_url=settings.external_api_url, - request_headers={"Content-Type": "application/json"}, - request_body=request_body, - response_status=response.status_code, - response_headers=dict(response.headers), - response_body=response.text, - retry_count=attempt - 1 - ) - - if response.status_code == 200: - logger.info(f"✅ 推送数据给DTC成功,状态码: {response.status_code}") - return True, attempt - 1 - else: - logger.warning(f"⚠️ DTC返回非成功状态码: {response.status_code}") - - # 如果是客户端错误(4xx),不重试 - if 400 <= response.status_code < 500: - logger.error(f"❌ 客户端错误,不重试: {response.status_code}") - return False, attempt - 1 - - except httpx.TimeoutException: - logger.warning(f"⏰ 推送数据给DTC超时,第{attempt}次尝试") - except httpx.RequestError as e: - logger.warning(f"🌐 推送数据给DTC请求错误,第{attempt}次尝试: {e}") - except Exception as e: - logger.error(f"❌ 推送数据给DTC异常,第{attempt}次尝试: {e}", exc_info=True) - - # 如果不是最后一次尝试,等待一段时间再重试 - if attempt < max_retries: - await asyncio.sleep(2 ** attempt) # 指数退避 - - # 所有重试都失败了 - logger.error(f"❌ 推送数据给DTC失败,已重试{max_retries}次") - return False, max_retries - - -async def mark_callback_log_completed(db: AsyncSession, callback_log_id: int): - """标记CallbackFailureLog记录为已完成""" - try: - # 更新指定日志记录为已完成 - result = await db.execute( - text(""" - UPDATE callback_failure_logs - SET is_completed = true - WHERE id = :log_id - AND is_completed = false - """), - {"log_id": callback_log_id} - ) - - updated_count = result.rowcount - await db.commit() - - logger.info(f"✅ 已标记回调日志为已完成,ID: {callback_log_id}, 更新记录数: {updated_count}") - - except Exception as e: - logger.error(f"❌ 标记回调日志为已完成失败: {e}", exc_info=True) - - -async def _log_dtc_push_call( - db: AsyncSession, - callback_failure_log_id: int, - request_url: str, - request_headers: dict, - request_body: dict, - response_status: int, - response_headers: dict, - response_body: str, - retry_count: int -): - """记录推送数据给DTC的日志""" - try: - api_log = ExternalApiLog( - callback_failure_log_id=callback_failure_log_id, - request_url=request_url, - request_headers=request_headers, - request_body=request_body, - response_status=response_status, - response_headers=response_headers, - response_body=response_body, - retry_count=retry_count - ) - - db.add(api_log) - await db.commit() - - logger.debug(f"📝 推送数据给DTC日志已记录,状态码: {response_status}") - - except Exception as e: - logger.error(f"❌ 记录推送数据给DTC日志失败: {e}", exc_info=True) \ No newline at end of file diff --git a/app/config.py b/app/config.py index 4984b52..81ffb62 100644 --- a/app/config.py +++ b/app/config.py @@ -2,6 +2,10 @@ from pydantic_settings import BaseSettings class Settings(BaseSettings): + @property + def redis_url(self) -> str: + """从Celery配置获取Redis URL""" + return self.celery_broker_url # 数据库配置 database_url: str = "" @@ -14,11 +18,21 @@ class Settings(BaseSettings): celery_timezone: str = "UTC" celery_enable_utc: bool = True + # Redis配置 (从Celery配置获取) + redis_password: str = "" + redis_max_connections: int = 20 + redis_timeout: int = 5 + # 业务配置 count_threshold: int = 3 # count阈值,大于等于此值直接返回 external_api_enabled: bool = False # 是否启用外部API调用 external_api_retry_max: int = 3 # 外部API最大重试次数 external_api_url: str = "" + + # Redis锁配置 (使用默认值) + redis_lock_timeout: int = 300 # 锁超时时间(秒) + redis_lock_max_retries: int = 10 # 最大重试次数 + redis_lock_retry_delay: float = 0.5 # 重试延迟(秒) # 日志配置 log_level: str = "INFO" # DEBUG, INFO, WARNING, ERROR, CRITICAL diff --git a/app/database.py b/app/database.py index 390c780..72669c5 100644 --- a/app/database.py +++ b/app/database.py @@ -65,7 +65,6 @@ AsyncSessionLocal = async_sessionmaker( engine, class_=AsyncSession, expire_on_commit=False ) - async def get_db(): async with AsyncSessionLocal() as session: try: @@ -73,7 +72,6 @@ async def get_db(): finally: await session.close() - async def init_db(): logger.info("📊 初始化数据库表结构...") try: diff --git a/app/external_api_processor.py b/app/external_api_processor.py deleted file mode 100644 index 0ca0ab9..0000000 --- a/app/external_api_processor.py +++ /dev/null @@ -1,287 +0,0 @@ -""" -异步调用外部接口的处理器 -""" - -from sqlalchemy.ext.asyncio import AsyncSession -from sqlalchemy import func, select -from fastapi import HTTPException -import httpx -import asyncio -from typing import Dict, Any, Optional -import json - -from app.database import CallbackFailureData, ExternalApiLog, CallbackFailureLog -from app.config import settings -from app.logger import get_logger - -logger = get_logger("external_api_processor") - - -async def get_callback_log_data( - db: AsyncSession, - callback_log_id: int -) -> tuple[bool, Optional[Dict[str, Any]], Optional[str]]: - """ - 从回调日志中获取请求头和请求体 - - Args: - db: 数据库会话 - callback_log_id: 回调日志ID - - Returns: - tuple[查询是否成功, 请求头字典, 请求体字符串] - """ - try: - # 查询回调日志 - query = select(CallbackFailureLog).where(CallbackFailureLog.id == callback_log_id) - result = await db.execute(query) - callback_log = result.scalar_one_or_none() - - if not callback_log: - logger.warning(f"⚠️ 未找到回调日志记录,ID: {callback_log_id}") - return False, None, None - - logger.info(f"✅ 成功获取回调日志,ID: {callback_log_id}") - - # 返回查询成功标识、请求头和请求体 - return True, callback_log.request_headers, callback_log.request_body - - except Exception as e: - logger.error(f"❌ 获取回调日志失败,ID: {callback_log_id}, 错误: {e}", exc_info=True) - return False, None, None - - -async def check_phone_number_threshold( - db: AsyncSession, - phone_number: str -) -> tuple[bool, bool]: - """ - 检查手机号在数据库中的出现次数是否超过阈值 - - Args: - db: 数据库会话 - phone_number: 要检查的手机号 - - Returns: - tuple[是否超过阈值, 是否查询成功] - """ - try: - # 查询手机号在数据库中出现的次数 - count_query = select(func.count(CallbackFailureData.id)).where( - CallbackFailureData.phone_number == phone_number - ) - result = await db.execute(count_query) - phone_count = result.scalar() or 0 - - logger.info(f"📊 手机号 {phone_number} 在数据库中出现次数: {phone_count}") - - # 判断是否超过阈值 - exceeds_threshold = phone_count >= settings.count_threshold - - if exceeds_threshold: - logger.info(f"✅ 手机号 {phone_number} 出现次数 {phone_count} >= {settings.count_threshold},超过阈值") - - return exceeds_threshold, True - - except Exception as e: - logger.error(f"❌ 查询手机号 {phone_number} 失败: {e}", exc_info=True) - # 查询失败时返回False,表示查询未成功 - return False, False - - -async def log_external_api_request( - db: AsyncSession, - callback_failure_log_id: int, - request_url: str, - request_headers: Dict[str, Any], - request_body: Dict[str, Any], - response_status: Optional[int] = None, - response_headers: Optional[Dict[str, Any]] = None, - response_body: Optional[str] = None, - retry_count: int = 0 -): - """记录外部API请求日志""" - try: - api_log = ExternalApiLog( - callback_failure_log_id=callback_failure_log_id, - request_url=request_url, - request_headers=request_headers, - request_body=request_body, - response_status=response_status, - response_headers=response_headers, - response_body=response_body, - retry_count=retry_count - ) - - db.add(api_log) - await db.commit() - - logger.info(f"✅ 外部API请求日志记录成功,ID: {api_log.id}") - - except Exception as e: - logger.error(f"❌ 记录外部API请求日志失败: {e}", exc_info=True) - - -async def call_external_api_with_retry( - db: AsyncSession, - request_body: Dict[str, Any], - max_retries: int = None, - callback_failure_log_id: int = None -) -> tuple[bool, int]: - """ - 调用外部API并支持重试机制 - - Args: - db: 数据库会话 - request_body: 请求体 - max_retries: 最大重试次数 - callback_failure_log_id: 回调失败日志ID - - Returns: - tuple[是否成功, 实际重试次数] - """ - if max_retries is None: - max_retries = settings.external_api_retry_max - - logger.info(f"🌐 开始调用外部API: {settings.external_api_url}, 最大重试次数: {max_retries}") - - headers = { - "Content-Type": "application/json", - "User-Agent": "AITalkCallbackService/1.0" - } - - for attempt in range(1, max_retries+1): - try: - logger.debug(f"📤 第{attempt + 1}次尝试调用外部API") - - async with httpx.AsyncClient(timeout=30.0) as client: - response = await client.post( - settings.external_api_url, - headers=headers, - json=request_body - ) - - logger.debug(f"📥 外部API响应: status={response.status_code}") - - # 记录每次尝试的结果 - await log_external_api_request( - db=db, - callback_failure_log_id=callback_failure_log_id, - request_url=settings.external_api_url, - request_headers=headers, - request_body=request_body, - response_status=response.status_code, - response_headers=dict(response.headers), - response_body=response.text, - retry_count=attempt - 1 - ) - - # 检查响应状态 - if response.status_code == 200: - logger.info(f"✅ 外部API调用成功,状态码: {response.status_code}") - return True, attempt - 1 - else: - logger.warning(f"⚠️ 外部API返回非成功状态码: {response.status_code}") - - # 如果是客户端错误(4xx),不重试 - if 400 <= response.status_code < 500: - logger.error(f"❌ 客户端错误,不重试: {response.status_code}") - return False, attempt - 1 - - except httpx.TimeoutException: - logger.warning(f"⏰ 外部API调用超时,第{attempt}次尝试") - except httpx.RequestError as e: - logger.warning(f"🌐 外部API请求错误,第{attempt}次尝试: {e}") - except Exception as e: - logger.error(f"❌ 外部API调用异常,第{attempt}次尝试: {e}", exc_info=True) - - # 如果不是最后一次尝试,等待一段时间再重试 - if attempt < max_retries: - await asyncio.sleep(2 ** attempt) # 指数退避 - - # 所有重试都失败了 - logger.error(f"❌ 外部API调用失败,已重试{max_retries}次") - return False, max_retries - - -async def process_external_api_call( - db: AsyncSession, - callback_log_id: int, - siteId: str, -): - """ - 异步调用外部接口 - - Args: - db: 数据库会话 - callback_log_id: 回调日志ID - siteId: 站点ID - - Returns: - None: 无返回值 - """ - - success, request_headers_json, request_body_json = await get_callback_log_data(db, callback_log_id) - if not success: - # 查询失败,处理错误情况 - logger.error(f"获取回调日志数据失败:{callback_log_id}") - return - - # 查询成功,可以使用获取到的数据 - logger.info(f"📋 请求头: {callback_log_id}|{request_headers_json}") - logger.info(f"📄 请求体: {callback_log_id}|{request_body_json}") - - request_headers = json.loads(request_headers_json) - request_body = json.loads(request_body_json) - - # 提取所有手机号并去重 - phone_numbers_set = set() - data_list = request_body.get('data', []) - for item in data_list: - number_data = item.get('number_data', {}) - phone_number = number_data.get('number') - if phone_number: - phone_numbers_set.add(phone_number) - - phone_numbers = list(phone_numbers_set) - logger.info(f"📱 提取到的手机号列表(去重后): {phone_numbers}") - - # 检查手机号列表是否为空 - if not phone_numbers or len(phone_numbers) == 0: - logger.warning(f"⚠️ 手机号列表为空,直接退出异步处理") - return - - # 检查去重后的手机号是否超过 settings.count_threshold,超过直接返回 - for phone_number in phone_numbers: - - # 检查手机号是否超过阈值 - exceeds_threshold, query_success = await check_phone_number_threshold(db, phone_number) - - # 如果查询失败,直接退出异步处理 - if not query_success: - logger.error(f"❌ 查询手机号 {phone_number} 失败,直接退出异步处理") - return - - if exceeds_threshold: - logger.info(f"✅ 手机号 {phone_number} 出现次数超过阈值,跳过处理") - continue - - # 手机号出现的次数少于settings.count_threshold,调用外部API - # 调用外部API并支持重试 - success, retry_count = await call_external_api_with_retry( - db=db, - request_body=request_body, - max_retries=settings.external_api_retry_max, - callback_failure_log_id=callback_log_id - ) - - if success: - logger.info(f"✅ 外部API调用成功,重试次数: {retry_count}") - return - else: - logger.error(f"❌ 外部API调用失败,已重试{retry_count}次") - return - - # 如果所有手机号都被跳过,返回成功但未处理 - logger.info(f"📝 所有手机号都被跳过,返回成功但未处理") - return \ No newline at end of file diff --git a/app/routes.py b/app/routes.py index b2b9f84..f846a01 100644 --- a/app/routes.py +++ b/app/routes.py @@ -1,99 +1,15 @@ from fastapi import APIRouter, Request, Body, HTTPException, Depends, Path from sqlalchemy.ext.asyncio import AsyncSession - - -from app.database import get_db, CallbackFailureLog +from app.database import get_db from app.models import CallbackResponse from app.config import settings from app.logger import get_logger -from app.celery_tasks import push_data_to_dtc_task +from app.callback_service import log_callback_request logger = get_logger("routes") router = APIRouter() - -async def log_callback_request( - db: AsyncSession, - request: Request, - site_id: str, - callback_data: dict -) -> bool: - """记录回调请求到数据库,返回操作是否成功""" - import json - - # 获取客户端IP地址 - client_ip = request.client.host if request.client else None - client_port = request.client.port if request.client else None - - # 获取服务器IP地址 - server_ip = None - server_port = None - if hasattr(request, 'scope') and 'server' in request.scope: - server_host, server_port_info = request.scope['server'] - server_ip = server_host - server_port = server_port_info - - # 准备请求头信息(直接记录原始请求头) - request_headers = dict(request.headers) - - # 准备请求体信息(直接记录原始请求体) - request_body = callback_data - - # 记录请求URL(JSON格式) - logger.info(f"🌐 请求URL: {request.url}") - - # 记录site_id(JSON格式) - logger.info(f"📝 site_id: {site_id}") - - # 记录server_ip(JSON格式) - server_info = { - "ip": server_ip, - "port": server_port - } - logger.info(f"🏠 server_ip: {server_info}") - - # 记录client_ip(JSON格式) - client_info = { - "ip": client_ip, - "port": client_port - } - logger.info(f"🖥️ client_ip: {client_info}") - - # 记录请求头(JSON格式) - logger.info(f"📋 请求头: {json.dumps(request_headers, ensure_ascii=False, indent=2)}") - - # 记录请求体(JSON格式) - logger.info(f"📄 请求体: {json.dumps(callback_data, ensure_ascii=False, indent=2)}") - - try: - # 保存到数据库 - callback_log = CallbackFailureLog( - site_id=site_id, - remote_address=f"{client_ip}:{client_port}" if client_ip and client_port else client_ip, - server_ip=f"{server_ip}:{server_port}" if server_ip and server_port else server_ip, - request_url=str(request.url), - request_headers=request_headers, # 保存原始请求头 - request_body=request_body # 保存从request.body获取的原始请求体 - ) - - db.add(callback_log) - await db.commit() - - logger.info(f"✅ 回调请求记录成功保存到数据库,ID: {callback_log.id}") - - # 返回成功标识 - return True - - except Exception as e: - logger.error(f"❌ 保存回调请求到数据库失败: {e}", exc_info=True) - # 返回失败标识 - return False - - - - - @router.post("/ai-talk/callback/{siteId}/failure", response_model=CallbackResponse) async def ai_talk_callback( request: Request, diff --git a/main.py b/main.py index 2eb8682..cc8832a 100644 --- a/main.py +++ b/main.py @@ -2,22 +2,69 @@ from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware from contextlib import asynccontextmanager from sqlalchemy import text +import redis.asyncio as redis +import threading +import subprocess +import sys +import os from app.config import settings from app.database import engine from app.routes import router from app.logger import get_logger, LoggerManager +from app.celery_app import celery_app # 初始化日志系统 LoggerManager.setup_logging() logger = get_logger("main") +# Celery Worker 进程管理 +celery_worker_process = None + + +def start_celery_worker(): + """启动 Celery Worker""" + global celery_worker_process + try: + logger.info("🌿 启动Celery Worker...") + + # 启动Celery worker + celery_app.start([ + 'worker', + '--loglevel=info', + '--concurrency=4', + '--prefetch-multiplier=1', + '--max-tasks-per-child=1000', + '--time-limit=300', # 5分钟任务超时 + '--soft-time-limit=240', # 4分钟软超时 + ]) + except Exception as e: + logger.error(f"❌ Celery Worker 启动失败: {e}") + + +def start_celery_beat(): + """启动 Celery Beat 调度器""" + try: + logger.info("📅 启动Celery Beat调度器...") + + # 启动Celery beat + celery_app.start([ + 'beat', + '--loglevel=info', + '--schedule=/tmp/celerybeat-schedule', + ]) + except Exception as e: + logger.error(f"❌ Celery Beat 启动失败: {e}") + @asynccontextmanager async def lifespan(app: FastAPI): # 启动时初始化 logger.info("🚀 应用启动中...") + + # Redis 连接对象 + redis_client = None try: # 验证数据库连接 @@ -35,6 +82,35 @@ async def lifespan(app: FastAPI): logger.error(f"❌ 数据库连接验证失败: {db_error}") raise + # 验证 Redis 连接 + logger.info("🔴 验证 Redis 连接...") + try: + redis_client = redis.from_url(settings.celery_broker_url) + # 执行 ping 命令验证连接 + await redis_client.ping() + logger.info("✅ Redis 连接验证成功") + + # 存储到应用状态中供其他组件使用 + app.state.redis_client = redis_client + + except Exception as redis_error: + logger.error(f"❌ Redis 连接验证失败: {redis_error}") + raise + + # 检查启动模式 + if len(sys.argv) > 1: + mode = sys.argv[1].replace("--mode=", "") + if mode == "all": + # 启动 Celery Worker (在后台线程中) + logger.info("🌿 启动Celery Worker...") + celery_worker_thread = threading.Thread(target=start_celery_worker, daemon=True) + celery_worker_thread.start() + + # 启动 Celery Beat 调度器 (在后台线程中) + logger.info("📅 启动Celery Beat调度器...") + celery_beat_thread = threading.Thread(target=start_celery_beat, daemon=True) + celery_beat_thread.start() + logger.info(f"🎉 {settings.app_name} 启动完成!") yield @@ -45,6 +121,15 @@ async def lifespan(app: FastAPI): finally: # 关闭时清理 logger.info("🛑 应用关闭中...") + + # 关闭 Redis 连接 + if redis_client: + try: + await redis_client.close() + logger.info("🔴 Redis 连接已关闭") + except Exception as e: + logger.warning(f"⚠️ 关闭 Redis 连接时出现警告: {e}") + logger.info("👋 应用已关闭") @@ -101,10 +186,46 @@ async def health_check(): raise HTTPException(status_code=404, detail="Not Found") logger.debug("💓 健康检查接口被访问") - return {"status": "healthy"} + + health_status = {"status": "healthy"} + + # 检查 Redis 连接状态 + try: + redis_client = getattr(app.state, 'redis_client', None) + if redis_client: + await redis_client.ping() + health_status["redis"] = "connected" + else: + health_status["redis"] = "disconnected" + except Exception as e: + health_status["redis"] = f"error: {str(e)}" + health_status["status"] = "degraded" + + return health_status if __name__ == "__main__": import uvicorn - - uvicorn.run("main:app", host="0.0.0.0", port=8000, reload=settings.debug) + import argparse + + parser = argparse.ArgumentParser(description="AI Talk Callback API") + parser.add_argument("--mode", choices=["api", "worker", "beat", "all"], + default="api", help="启动模式: api(仅API), worker(仅Celery Worker), beat(仅Celery Beat), all(全部)") + args = parser.parse_args() + + if args.mode == "api": + # 仅启动 FastAPI 应用 + logger.info("🚀 启动FastAPI应用...") + uvicorn.run("main:app", host="0.0.0.0", port=8000, reload=settings.debug) + elif args.mode == "worker": + # 仅启动 Celery Worker + logger.info("🌿 启动Celery Worker...") + start_celery_worker() + elif args.mode == "beat": + # 仅启动 Celery Beat + logger.info("📅 启动Celery Beat调度器...") + start_celery_beat() + elif args.mode == "all": + # 启动 FastAPI + Celery Worker + Celery Beat + logger.info("🚀 启动完整服务栈(FastAPI + Celery Worker + Celery Beat)...") + uvicorn.run("main:app", host="0.0.0.0", port=8000, reload=settings.debug) diff --git a/test_callback.py b/test_callback.py index ffc9632..6cef057 100644 --- a/test_callback.py +++ b/test_callback.py @@ -88,7 +88,7 @@ class TestAITalkCallback: @patch('app.routes.get_db') @patch('app.routes.redis_manager') - @patch('app.routes.log_callback_request') + @patch('app.callback_service.log_callback_request') def test_callback_count_below_threshold_success( self, mock_log_callback, @@ -126,7 +126,7 @@ class TestAITalkCallback: @patch('app.routes.get_db') @patch('app.routes.redis_manager') - @patch('app.routes.log_callback_request') + @patch('app.callback_service.log_callback_request') def test_callback_count_above_threshold_direct_return( self, mock_log_callback, @@ -154,7 +154,7 @@ class TestAITalkCallback: @patch('app.routes.get_db') @patch('app.routes.redis_manager') - @patch('app.routes.log_callback_request') + @patch('app.callback_service.log_callback_request') def test_callback_external_api_failure( self, mock_log_callback, @@ -215,7 +215,7 @@ class TestAITalkCallback: @patch('app.routes.get_db') @patch('app.routes.redis_manager') - @patch('app.routes.log_callback_request') + @patch('app.callback_service.log_callback_request') def test_callback_empty_data_list( self, mock_log_callback, @@ -251,7 +251,7 @@ class TestAITalkCallback: @patch('app.routes.get_db') @patch('app.routes.redis_manager') - @patch('app.routes.log_callback_request') + @patch('app.callback_service.log_callback_request') def test_callback_database_error( self, mock_log_callback, @@ -273,7 +273,7 @@ class TestAITalkCallback: @patch('app.routes.get_db') @patch('app.routes.redis_manager') - @patch('app.routes.log_callback_request') + @patch('app.callback_service.log_callback_request') def test_callback_redis_lock_error( self, mock_log_callback,