推送数据给DTC改为独立执行的任务

This commit is contained in:
mark.tian
2025-12-04 15:05:56 +08:00
parent 929f00b34e
commit d3bbe5f649
9 changed files with 503 additions and 96 deletions

View File

@@ -1,6 +1,10 @@
# 数据库配置 # 数据库配置
DATABASE_URL=postgresql+asyncpg://user:password@localhost:5432/ai_talk_callback_db 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 COUNT_THRESHOLD=3
EXTERNAL_API_ENABLED=false EXTERNAL_API_ENABLED=false

32
app/celery_app.py Normal file
View File

@@ -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应用配置完成")

410
app/celery_tasks.py Normal file
View File

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

View File

@@ -5,6 +5,15 @@ class Settings(BaseSettings):
# 数据库配置 # 数据库配置
database_url: str = "" 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阈值,大于等于此值直接返回 count_threshold: int = 3 # count阈值,大于等于此值直接返回
external_api_enabled: bool = False # 是否启用外部API调用 external_api_enabled: bool = False # 是否启用外部API调用

View File

@@ -1,7 +1,7 @@
from datetime import datetime from datetime import datetime
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
from sqlalchemy.orm import DeclarativeBase 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.config import settings
from app.logger import get_logger from app.logger import get_logger
from sqlalchemy.sql.expression import func from sqlalchemy.sql.expression import func
@@ -23,6 +23,7 @@ class CallbackFailureLog(Base):
request_url = Column(String(500), nullable=False, comment="请求URL") request_url = Column(String(500), nullable=False, comment="请求URL")
request_headers = Column(JSON, nullable=False, comment="请求头") request_headers = Column(JSON, nullable=False, comment="请求头")
request_body = 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="创建时间") created_at = Column(DateTime, default=datetime.now(), comment="创建时间")

View File

@@ -1,12 +1,12 @@
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 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.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.external_api_processor import process_external_api_call from app.celery_tasks import push_data_to_dtc_task
logger = get_logger("routes") logger = get_logger("routes")
@@ -18,8 +18,8 @@ async def log_callback_request(
request: Request, request: Request,
site_id: str, site_id: str,
callback_data: dict callback_data: dict
) -> tuple[bool, Optional[int]]: ) -> bool:
"""记录回调请求到数据库,返回操作是否成功和记录ID""" """记录回调请求到数据库,返回操作是否成功"""
import json import json
# 获取客户端IP地址 # 获取客户端IP地址
@@ -82,82 +82,16 @@ async def log_callback_request(
logger.info(f"✅ 回调请求记录成功保存到数据库,ID: {callback_log.id}") logger.info(f"✅ 回调请求记录成功保存到数据库,ID: {callback_log.id}")
# 返回成功标识和记录ID # 返回成功标识
return True, callback_log.id return True
except Exception as e: except Exception as e:
logger.error(f"❌ 保存回调请求到数据库失败: {e}", exc_info=True) 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) @router.post("/ai-talk/callback/{siteId}/failure", response_model=CallbackResponse)
@@ -172,27 +106,12 @@ async def ai_talk_callback(
""" """
try: try:
# 记录回调请求(包含siteId) # 记录回调请求(包含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: if not success:
logger.error("❌ 回调请求日志保存失败") logger.error("❌ 回调请求日志保存失败")
raise HTTPException(status_code=500, detail="回调请求日志保存失败") 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( return CallbackResponse(
success=True, success=True,
@@ -202,7 +121,6 @@ async def ai_talk_callback(
site_id=siteId site_id=siteId
) )
except HTTPException: except HTTPException:
raise raise
except Exception as e: except Exception as e:

28
celery_worker.py Normal file
View File

@@ -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分钟软超时
])

View File

@@ -3,6 +3,8 @@ uvicorn[standard]>=0.38.0
sqlalchemy>=2.0.44 sqlalchemy>=2.0.44
asyncpg>=0.13.0 asyncpg>=0.13.0
alembic>=1.17.2 alembic>=1.17.2
celery>=5.3.0
redis>=4.5.0
pydantic>=2.12.5 pydantic>=2.12.5
pydantic-settings>=2.12.0 pydantic-settings>=2.12.0
python-multipart>=0.0.20 python-multipart>=0.0.20

View File

@@ -139,6 +139,7 @@ async def main():
print("🚀 开始数据库初始化...") print("🚀 开始数据库初始化...")
print(f"📋 配置信息:") print(f"📋 配置信息:")
print(f" - 数据库URL: {settings.database_url}") print(f" - 数据库URL: {settings.database_url}")
print(f" - Celery Broker URL: {settings.celery_broker_url}")
print(f" - 应用名称: {settings.app_name}") print(f" - 应用名称: {settings.app_name}")
print(f" - 外部API URL: {settings.external_api_url}") print(f" - 外部API URL: {settings.external_api_url}")
print() print()
@@ -163,9 +164,11 @@ async def main():
print("🎉 数据库初始化完成!") print("🎉 数据库初始化完成!")
print("\n📝 下一步:") print("\n📝 下一步:")
print(" 1. 配置 .env 文件中的数据库连接信息") print(" 1. 配置 .env 文件中的数据库连接信息")
print(" 2. 运行应用: python main.py") print(" 2. 确保 Redis 服务正在运行(用于Celery)")
print(f" 3. 访问API文档: http://localhost:8000/docs") print(" 3. 启动 Celery Worker: python celery_worker.py")
print(f" 4. 接口地址: POST /ai-talk/callback/{{siteId}}/failure") 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__": if __name__ == "__main__":