diff --git a/.env.example b/.env.example index 3274ca0..4589b99 100644 --- a/.env.example +++ b/.env.example @@ -1,6 +1,10 @@ # 数据库配置 DATABASE_URL=postgresql+asyncpg://user:password@localhost:5432/ai_talk_callback_db +# Celery配置 +CELERY_BROKER_URL=redis://localhost:6379/0 +CELERY_RESULT_BACKEND=redis://localhost:6379/0 + # 业务配置 COUNT_THRESHOLD=3 EXTERNAL_API_ENABLED=false diff --git a/app/celery_app.py b/app/celery_app.py new file mode 100644 index 0000000..13bce2b --- /dev/null +++ b/app/celery_app.py @@ -0,0 +1,32 @@ +""" +Celery应用配置 +""" +from celery import Celery +from app.config import settings +from app.logger import get_logger + +logger = get_logger("celery") + +# 创建Celery应用实例 +celery_app = Celery( + "ai_talk_callback", + broker=settings.celery_broker_url, + backend=settings.celery_result_backend, + include=['app.celery_tasks'] +) + +# Celery配置 +celery_app.conf.update( + task_serializer=settings.celery_task_serializer, + result_serializer=settings.celery_result_serializer, + accept_content=settings.celery_accept_content, + timezone=settings.celery_timezone, + enable_utc=settings.celery_enable_utc, + task_track_started=True, + task_time_limit=30 * 60, # 30分钟超时 + task_soft_time_limit=25 * 60, # 25分钟软超时 + worker_prefetch_multiplier=1, + worker_max_tasks_per_child=1000, +) + +logger.info("🌿 Celery应用配置完成") \ No newline at end of file diff --git a/app/celery_tasks.py b/app/celery_tasks.py new file mode 100644 index 0000000..6f39ea4 --- /dev/null +++ b/app/celery_tasks.py @@ -0,0 +1,410 @@ +""" +Celery任务定义 +""" +from celery import current_task +from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker +from sqlalchemy import 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 + +logger = get_logger("celery_tasks") + +# 创建独立的数据库连接用于Celery任务 +engine = create_async_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任务 + 自动获取一条未完成的回调请求进行处理 + """ + logger.info(f"🌿 开始推送数据给DTC任务") + + # 使用asyncio运行异步逻辑 + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + + 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: + 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: + # 更新任务状态 + current_task.update_state( + state='PROGRESS', + meta={'current': processed_count + skipped_count, '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 + + # 推送数据给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 + ) + + if success: + logger.info(f"✅ 推送数据给DTC成功,手机号: {phone_number}, 重试次数: {retry_count}") + processed_count += 1 + else: + logger.error(f"❌ 推送数据给DTC失败,手机号: {phone_number}, 重试次数: {retry_count}") + skipped_count += 1 + + 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) + } + + 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 + + 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) + # 不重新抛出异常,避免影响主业务流程 + + +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 0bceb1e..4984b52 100644 --- a/app/config.py +++ b/app/config.py @@ -5,6 +5,15 @@ class Settings(BaseSettings): # 数据库配置 database_url: str = "" + # Celery配置 + celery_broker_url: str = "redis://localhost:6379/0" + celery_result_backend: str = "redis://localhost:6379/0" + celery_task_serializer: str = "json" + celery_result_serializer: str = "json" + celery_accept_content: list = ["json"] + celery_timezone: str = "UTC" + celery_enable_utc: bool = True + # 业务配置 count_threshold: int = 3 # count阈值,大于等于此值直接返回 external_api_enabled: bool = False # 是否启用外部API调用 diff --git a/app/database.py b/app/database.py index 8d4eafe..390c780 100644 --- a/app/database.py +++ b/app/database.py @@ -1,7 +1,7 @@ from datetime import datetime from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker from sqlalchemy.orm import DeclarativeBase -from sqlalchemy import Column, String, Integer, DateTime, Text, JSON, Index, text +from sqlalchemy import Column, String, Integer, DateTime, Text, JSON, Index, text, Boolean from app.config import settings from app.logger import get_logger from sqlalchemy.sql.expression import func @@ -23,6 +23,7 @@ class CallbackFailureLog(Base): request_url = Column(String(500), nullable=False, comment="请求URL") request_headers = Column(JSON, nullable=False, comment="请求头") request_body = Column(JSON, nullable=False, comment="请求体") + is_completed = Column(Boolean, default=False, comment="是否已完成") created_at = Column(DateTime, default=datetime.now(), comment="创建时间") diff --git a/app/routes.py b/app/routes.py index c072ecd..b2b9f84 100644 --- a/app/routes.py +++ b/app/routes.py @@ -1,12 +1,12 @@ from fastapi import APIRouter, Request, Body, HTTPException, Depends, Path from sqlalchemy.ext.asyncio import AsyncSession -from typing import Optional -from app.database import get_db, CallbackFailureLog, CallbackFailureData + +from app.database import get_db, CallbackFailureLog from app.models import CallbackResponse from app.config import settings from app.logger import get_logger -from app.external_api_processor import process_external_api_call +from app.celery_tasks import push_data_to_dtc_task logger = get_logger("routes") @@ -18,8 +18,8 @@ async def log_callback_request( request: Request, site_id: str, callback_data: dict -) -> tuple[bool, Optional[int]]: - """记录回调请求到数据库,返回操作是否成功和记录ID""" +) -> bool: + """记录回调请求到数据库,返回操作是否成功""" import json # 获取客户端IP地址 @@ -82,82 +82,16 @@ async def log_callback_request( logger.info(f"✅ 回调请求记录成功保存到数据库,ID: {callback_log.id}") - # 返回成功标识和记录ID - return True, callback_log.id + # 返回成功标识 + return True except Exception as e: logger.error(f"❌ 保存回调请求到数据库失败: {e}", exc_info=True) # 返回失败标识 - return False, None + return False + -async def save_callback_data_items( - db: AsyncSession, - callback_data_items: list, - callback_failure_log_id: int -): - """保存callback_data.data中的数据到数据库""" - import json - from datetime import datetime - - 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) - # 不重新抛出异常,避免影响主业务流程 @router.post("/ai-talk/callback/{siteId}/failure", response_model=CallbackResponse) @@ -172,27 +106,12 @@ async def ai_talk_callback( """ try: # 记录回调请求(包含siteId) - success, callback_log_id = await log_callback_request(db, request, siteId, callback_data) + success = await log_callback_request(db, request, siteId, callback_data) if not success: logger.error("❌ 回调请求日志保存失败") raise HTTPException(status_code=500, detail="回调请求日志保存失败") - if callback_log_id: - # 保存callback_data.data中的数据 - data_list = callback_data.get('data', []) - if data_list and len(data_list) > 0: - await save_callback_data_items(db, data_list, callback_log_id) - - # 检查是否启用外部API调用 - if settings.external_api_enabled: - # 异步调用外部接口,不等待结果(只有在成功获取callback_log_id时才调用) - if success and callback_log_id: - import asyncio - asyncio.create_task(process_external_api_call(db, callback_log_id, siteId)) - else: - logger.info(f"🔌 外部API调用已禁用,直接返回成功") - # 立即返回成功响应 return CallbackResponse( success=True, @@ -201,7 +120,6 @@ async def ai_talk_callback( retry_count=0, site_id=siteId ) - except HTTPException: raise diff --git a/celery_worker.py b/celery_worker.py new file mode 100644 index 0000000..afc742a --- /dev/null +++ b/celery_worker.py @@ -0,0 +1,28 @@ +#!/usr/bin/env python3 +""" +Celery Worker 启动脚本 +""" +import os +import sys + +# 添加项目根目录到Python路径 +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + +from app.celery_app import celery_app +from app.logger import get_logger + +logger = get_logger("celery_worker") + +if __name__ == "__main__": + 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分钟软超时 + ]) \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index ea3b11d..264e742 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,6 +3,8 @@ uvicorn[standard]>=0.38.0 sqlalchemy>=2.0.44 asyncpg>=0.13.0 alembic>=1.17.2 +celery>=5.3.0 +redis>=4.5.0 pydantic>=2.12.5 pydantic-settings>=2.12.0 python-multipart>=0.0.20 diff --git a/setup.py b/setup.py index 608bff3..17e0390 100644 --- a/setup.py +++ b/setup.py @@ -139,6 +139,7 @@ async def main(): print("🚀 开始数据库初始化...") print(f"📋 配置信息:") print(f" - 数据库URL: {settings.database_url}") + print(f" - Celery Broker URL: {settings.celery_broker_url}") print(f" - 应用名称: {settings.app_name}") print(f" - 外部API URL: {settings.external_api_url}") print() @@ -163,9 +164,11 @@ async def main(): print("🎉 数据库初始化完成!") print("\n📝 下一步:") print(" 1. 配置 .env 文件中的数据库连接信息") - print(" 2. 运行应用: python main.py") - print(f" 3. 访问API文档: http://localhost:8000/docs") - print(f" 4. 接口地址: POST /ai-talk/callback/{{siteId}}/failure") + print(" 2. 确保 Redis 服务正在运行(用于Celery)") + print(" 3. 启动 Celery Worker: python celery_worker.py") + print(" 4. 运行应用: python main.py") + print(f" 5. 访问API文档: http://localhost:8000/docs") + print(f" 6. 接口地址: POST /ai-talk/callback/{{siteId}}/failure") if __name__ == "__main__":