优化推送任务

This commit is contained in:
mark.tian
2025-12-05 06:20:14 +08:00
parent d3bbe5f649
commit 038363a17e
8 changed files with 288 additions and 751 deletions

View File

@@ -5,6 +5,16 @@ DATABASE_URL=postgresql+asyncpg://user:password@localhost:5432/ai_talk_callback_
CELERY_BROKER_URL=redis://localhost:6379/0 CELERY_BROKER_URL=redis://localhost:6379/0
CELERY_RESULT_BACKEND=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 COUNT_THRESHOLD=3
EXTERNAL_API_ENABLED=false EXTERNAL_API_ENABLED=false

View File

@@ -2,409 +2,174 @@
Celery任务定义 Celery任务定义
""" """
from celery import current_task from celery import current_task
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker from sqlalchemy import create_engine, text
from sqlalchemy import text
import json import json
import httpx
import asyncio
from app.celery_app import celery_app from app.celery_app import celery_app
from app.config import settings from app.config import settings
from app.logger import get_logger from app.logger import get_logger
from app.database import CallbackFailureLog, ExternalApiLog, CallbackFailureData 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") logger = get_logger("celery_tasks")
# 创建独立的数据库连接用于Celery任务 # 创建同步数据库连接用于Celery任务
engine = create_async_engine( engine = create_engine(
settings.database_url, settings.database_url,
echo=settings.debug, echo=settings.debug,
future=True 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') @celery_app.task(bind=True, name='push_data_to_dtc')
def push_data_to_dtc_task(self): def push_data_to_dtc_task(self):
""" """
推送数据给DTC的Celery任务 推送数据给DTC的Celery任务
自动获取一条未完成的回调请求进行处理 自动获取一条未完成的回调请求进行处理
""" """
task_name = 'push_data_to_dtc'
logger.info(f"🌿 开始推送数据给DTC任务") logger.info(f"🌿 开始推送数据给DTC任务")
# 使用asyncio运行异步逻辑 # 获取分布式锁,使用任务名称作为锁标识
loop = asyncio.new_event_loop() lock = redis_manager.create_lock(f"celery_task:{task_name}", timeout=300) # 5分钟超时
asyncio.set_event_loop(loop)
try: try:
result = loop.run_until_complete( # 尝试获取锁
_push_data_to_dtc_async(self.request.id) if not lock.acquire(blocking=False):
) logger.warning(f"⚠️ 任务 {task_name} 正在执行中,跳过本次执行")
return result return {"status": "skipped", "message": f"任务 {task_name} 正在执行中,跳过本次执行"}
except Exception as e:
logger.error(f"❌ 推送数据给DTC任务执行失败: {e}", exc_info=True)
raise
finally:
loop.close()
logger.info(f"🔒 成功获取任务 {task_name} 的分布式锁")
async def _push_data_to_dtc_async(task_id: str):
"""异步推送数据给DTC的核心逻辑"""
async with AsyncSessionLocal() as db:
try: try:
# 获取一条未完成的回调请求(按创建时间取最小值) with engine.connect() as conn:
callback_log_data = await get_uncompleted_callback_log(db) # 获取一条未完成的回调请求(按创建时间取最小值)
if not callback_log_data: query_success, callback_log_data = get_uncompleted_callback_log(conn)
logger.info("📋 没有找到未完成的回调请求") if not query_success:
return {"status": "skipped", "message": "没有找到未完成的回调请求"} logger.error("❌ 查询未完成的回调请求失败,尝试再查一次")
query_success, callback_log_data = get_uncompleted_callback_log(conn)
if not query_success:
logger.error("❌ 再次查询未完成的回调请求失败,停止任务执行")
raise Exception("查询未完成的回调请求失败,任务停止执行")
callback_log_id, site_id, request_headers_json, request_body_json = callback_log_data if not callback_log_data:
logger.info(f"📋 获取到未完成的回调请求: ID={callback_log_id}, site_id={site_id}") logger.info("📋 没有找到未完成的回调请求")
return {"status": "skipped", "message": "没有找到未完成的回调请求"}
# 解析请求头和请求体 callback_log_id, site_id, request_headers_json, request_body_json = callback_log_data
request_headers = json.loads(request_headers_json) logger.info(f"📋 获取到未完成的回调请求: ID={callback_log_id}, site_id={site_id}")
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( current_task.update_state(
state='PROGRESS', 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) request_headers = json.loads(request_headers_json)
request_body = json.loads(request_body_json)
if not query_success: # 保存callback_data.data中的数据
logger.error(f"❌ 查询手机号 {phone_number} 失败,跳过处理") data_list = request_body.get('data', [])
skipped_count += 1 if data_list and len(data_list) > 0:
continue 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}失败,任务停止执行")
if exceeds_threshold: # 提取所有手机号并去重
logger.info(f"✅ 手机号 {phone_number} 出现次数超过阈值,跳过处理") phone_numbers_set = set()
skipped_count += 1 data_list = request_body.get('data', [])
continue 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)
# 推送数据给DTC phone_numbers = list(phone_numbers_set)
success, retry_count = await push_data_to_dtc_with_retry( logger.info(f"📱 提取到的手机号列表(去重后): {phone_numbers}")
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}") valid_phone_numbers = []
processed_count += 1 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: else:
logger.error(f"❌ 推送数据给DTC失败,手机号: {phone_number}, 重试次数: {retry_count}") logger.info("📋 没有有效的手机号需要处理")
skipped_count += 1
logger.info(f"🎉 推送数据给DTC处理完成,成功: {processed_count}, 跳过: {skipped_count}") logger.info(f"🎉 推送数据给DTC处理完成,成功: {processed_count}, 跳过: {skipped_count}")
# 标记CallbackFailureLog为已完成 # 标记CallbackFailureLog为已完成
await mark_callback_log_completed(db, callback_log_id) mark_callback_log_completed(conn, callback_log_id)
return { return {
"status": "completed", "status": "completed",
"processed": processed_count, "processed": processed_count,
"skipped": skipped_count, "skipped": skipped_count,
"total": len(phone_numbers) "total": len(phone_numbers)
} }
except Exception as e: except Exception as e:
logger.error(f"❌ 推送数据给DTC时发生错误: {e}", exc_info=True) logger.error(f"❌ 推送数据给DTC时发生错误: {e}", exc_info=True)
return {"status": "error", "message": str(e)} return {"status": "error", "message": str(e)}
finally:
async def save_callback_data_items( # 释放分布式锁
db: AsyncSession, try:
callback_data_items: list, lock.release()
callback_failure_log_id: int logger.info(f"🔓 释放任务 {task_name} 的分布式锁")
): except Exception as e:
"""保存callback_data.data中的数据到数据库""" logger.error(f"❌ 释放任务 {task_name} 的分布式锁失败: {e}")
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: 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)

View File

@@ -2,6 +2,10 @@ from pydantic_settings import BaseSettings
class Settings(BaseSettings): class Settings(BaseSettings):
@property
def redis_url(self) -> str:
"""从Celery配置获取Redis URL"""
return self.celery_broker_url
# 数据库配置 # 数据库配置
database_url: str = "" database_url: str = ""
@@ -14,12 +18,22 @@ class Settings(BaseSettings):
celery_timezone: str = "UTC" celery_timezone: str = "UTC"
celery_enable_utc: bool = True celery_enable_utc: bool = True
# Redis配置 (从Celery配置获取)
redis_password: str = ""
redis_max_connections: int = 20
redis_timeout: int = 5
# 业务配置 # 业务配置
count_threshold: int = 3 # count阈值,大于等于此值直接返回 count_threshold: int = 3 # count阈值,大于等于此值直接返回
external_api_enabled: bool = False # 是否启用外部API调用 external_api_enabled: bool = False # 是否启用外部API调用
external_api_retry_max: int = 3 # 外部API最大重试次数 external_api_retry_max: int = 3 # 外部API最大重试次数
external_api_url: str = "" 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 log_level: str = "INFO" # DEBUG, INFO, WARNING, ERROR, CRITICAL
log_file: str = "logs/app.log" # 日志文件路径 log_file: str = "logs/app.log" # 日志文件路径

View File

@@ -65,7 +65,6 @@ AsyncSessionLocal = async_sessionmaker(
engine, class_=AsyncSession, expire_on_commit=False engine, class_=AsyncSession, expire_on_commit=False
) )
async def get_db(): async def get_db():
async with AsyncSessionLocal() as session: async with AsyncSessionLocal() as session:
try: try:
@@ -73,7 +72,6 @@ async def get_db():
finally: finally:
await session.close() await session.close()
async def init_db(): async def init_db():
logger.info("📊 初始化数据库表结构...") logger.info("📊 初始化数据库表结构...")
try: try:

View File

@@ -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

View File

@@ -1,99 +1,15 @@
from fastapi import APIRouter, Request, Body, HTTPException, Depends, Path from fastapi import APIRouter, Request, Body, HTTPException, Depends, Path
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from app.database import get_db, CallbackFailureLog
from app.models import CallbackResponse from app.models import CallbackResponse
from app.config import settings from app.config import settings
from app.logger import get_logger 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") logger = get_logger("routes")
router = APIRouter() 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) @router.post("/ai-talk/callback/{siteId}/failure", response_model=CallbackResponse)
async def ai_talk_callback( async def ai_talk_callback(
request: Request, request: Request,

125
main.py
View File

@@ -2,23 +2,70 @@ from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from sqlalchemy import text from sqlalchemy import text
import redis.asyncio as redis
import threading
import subprocess
import sys
import os
from app.config import settings from app.config import settings
from app.database import engine from app.database import engine
from app.routes import router from app.routes import router
from app.logger import get_logger, LoggerManager from app.logger import get_logger, LoggerManager
from app.celery_app import celery_app
# 初始化日志系统 # 初始化日志系统
LoggerManager.setup_logging() LoggerManager.setup_logging()
logger = get_logger("main") 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 @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI):
# 启动时初始化 # 启动时初始化
logger.info("🚀 应用启动中...") logger.info("🚀 应用启动中...")
# Redis 连接对象
redis_client = None
try: try:
# 验证数据库连接 # 验证数据库连接
logger.info("📊 验证数据库连接...") logger.info("📊 验证数据库连接...")
@@ -35,6 +82,35 @@ async def lifespan(app: FastAPI):
logger.error(f"❌ 数据库连接验证失败: {db_error}") logger.error(f"❌ 数据库连接验证失败: {db_error}")
raise 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} 启动完成!") logger.info(f"🎉 {settings.app_name} 启动完成!")
yield yield
@@ -45,6 +121,15 @@ async def lifespan(app: FastAPI):
finally: finally:
# 关闭时清理 # 关闭时清理
logger.info("🛑 应用关闭中...") 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("👋 应用已关闭") logger.info("👋 应用已关闭")
@@ -101,10 +186,46 @@ async def health_check():
raise HTTPException(status_code=404, detail="Not Found") raise HTTPException(status_code=404, detail="Not Found")
logger.debug("💓 健康检查接口被访问") 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__": if __name__ == "__main__":
import uvicorn import uvicorn
import argparse
uvicorn.run("main:app", host="0.0.0.0", port=8000, reload=settings.debug) 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)

View File

@@ -88,7 +88,7 @@ class TestAITalkCallback:
@patch('app.routes.get_db') @patch('app.routes.get_db')
@patch('app.routes.redis_manager') @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( def test_callback_count_below_threshold_success(
self, self,
mock_log_callback, mock_log_callback,
@@ -126,7 +126,7 @@ class TestAITalkCallback:
@patch('app.routes.get_db') @patch('app.routes.get_db')
@patch('app.routes.redis_manager') @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( def test_callback_count_above_threshold_direct_return(
self, self,
mock_log_callback, mock_log_callback,
@@ -154,7 +154,7 @@ class TestAITalkCallback:
@patch('app.routes.get_db') @patch('app.routes.get_db')
@patch('app.routes.redis_manager') @patch('app.routes.redis_manager')
@patch('app.routes.log_callback_request') @patch('app.callback_service.log_callback_request')
def test_callback_external_api_failure( def test_callback_external_api_failure(
self, self,
mock_log_callback, mock_log_callback,
@@ -215,7 +215,7 @@ class TestAITalkCallback:
@patch('app.routes.get_db') @patch('app.routes.get_db')
@patch('app.routes.redis_manager') @patch('app.routes.redis_manager')
@patch('app.routes.log_callback_request') @patch('app.callback_service.log_callback_request')
def test_callback_empty_data_list( def test_callback_empty_data_list(
self, self,
mock_log_callback, mock_log_callback,
@@ -251,7 +251,7 @@ class TestAITalkCallback:
@patch('app.routes.get_db') @patch('app.routes.get_db')
@patch('app.routes.redis_manager') @patch('app.routes.redis_manager')
@patch('app.routes.log_callback_request') @patch('app.callback_service.log_callback_request')
def test_callback_database_error( def test_callback_database_error(
self, self,
mock_log_callback, mock_log_callback,
@@ -273,7 +273,7 @@ class TestAITalkCallback:
@patch('app.routes.get_db') @patch('app.routes.get_db')
@patch('app.routes.redis_manager') @patch('app.routes.redis_manager')
@patch('app.routes.log_callback_request') @patch('app.callback_service.log_callback_request')
def test_callback_redis_lock_error( def test_callback_redis_lock_error(
self, self,
mock_log_callback, mock_log_callback,