celery任务改为异步访问数据库

This commit is contained in:
mark.tian
2025-12-10 11:40:56 +08:00
parent d3d3502f5a
commit 572cc3113f
5 changed files with 456 additions and 213 deletions

9
2.0.0 Normal file
View File

@@ -0,0 +1,9 @@
Looking in indexes: https://mirrors.aliyun.com/pypi/simple/, https://pypi.tuna.tsinghua.edu.cn/simple
Collecting aioredis
Using cached https://pypi.tuna.tsinghua.edu.cn/packages/9b/a9/0da089c3ae7a31cbcd2dcf0214f6f571e1295d292b6139e2bac68ec081d0/aioredis-2.0.1-py3-none-any.whl (71 kB)
Collecting async-timeout (from aioredis)
Downloading https://pypi.tuna.tsinghua.edu.cn/packages/fe/ba/e2081de779ca30d473f21f5b30e0e737c438205440784c7dfc81efc2b029/async_timeout-5.0.1-py3-none-any.whl (6.2 kB)
Requirement already satisfied: typing-extensions in e:\ai\code\ai-talk-callback\.venv\lib\site-packages (from aioredis) (4.15.0)
Installing collected packages: async-timeout, aioredis
Successfully installed aioredis-2.0.1 async-timeout-5.0.1

View File

@@ -3,11 +3,10 @@
包含与回调处理相关的业务逻辑函数
"""
import json
import time
import asyncio
import httpx
from datetime import datetime
from typing import Optional, Tuple, Dict, Any
from sqlalchemy import text
from fastapi import Request
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
@@ -95,8 +94,8 @@ async def log_callback_request(
return False
def save_callback_data_items(
conn,
async def save_callback_data_items(
db: AsyncSession,
callback_data_items: list,
callback_failure_log_id: int
) -> bool:
@@ -105,12 +104,14 @@ def save_callback_data_items(
Returns:
bool: 保存是否成功,True表示成功,False表示失败或跳过保存
"""
from sqlalchemy import text
try:
logger.info(f"📊 共传入callback_log_id {callback_failure_log_id} 的 {len(callback_data_items)} 条数据")
# 检查是否已经存在该callback_log_id的数据
existing_data_result = conn.execute(
existing_data_result = await db.execute(
text("SELECT * FROM callback_failure_data WHERE callback_failure_log_id = :log_id"),
{"log_id": callback_failure_log_id}
)
@@ -157,7 +158,7 @@ def save_callback_data_items(
raw_data_json = json.dumps(item, ensure_ascii=False)
# 直接插入数据库
conn.execute(
await db.execute(
text("""
INSERT INTO callback_failure_data
(callback_failure_log_id, phone_number, task_id, user_id, status, status_description, raw_data, calldate)
@@ -175,17 +176,18 @@ def save_callback_data_items(
}
)
conn.commit()
await db.commit()
logger.info(f"✅ 成功保存 {len(callback_data_items)} 条callback_data记录到数据库")
return True
except Exception as e:
logger.error(f"❌ 保存callback_data到数据库失败: {e}", exc_info=True)
await db.rollback()
# 不重新抛出异常,避免影响主业务流程
return False
def get_uncompleted_callback_log(conn) -> Tuple[bool, Optional[Tuple[int, str, str, str]]]:
async def get_uncompleted_callback_log(db: AsyncSession) -> Tuple[bool, Optional[Tuple[int, str, str, str]]]:
"""获取一条未完成的回调请求(按创建时间取最小值)
Returns:
@@ -193,9 +195,11 @@ def get_uncompleted_callback_log(conn) -> Tuple[bool, Optional[Tuple[int, str, s
- 第一个值表示查询是否成功(True表示成功,False表示失败)
- 第二个值为回调日志数据元组或None
"""
from sqlalchemy import text
try:
# 查询一条未完成的回调日志(按创建时间升序排列,取第一条)
result = conn.execute(
result = await db.execute(
text("""
SELECT id, site_id, request_headers, request_body
FROM callback_failure_logs
@@ -221,11 +225,13 @@ def get_uncompleted_callback_log(conn) -> Tuple[bool, Optional[Tuple[int, str, s
return False, None
def get_callback_log_data(conn, callback_log_id: int) -> Tuple[bool, Optional[str], Optional[str]]:
async def get_callback_log_data(db: AsyncSession, callback_log_id: int) -> Tuple[bool, Optional[str], Optional[str]]:
"""获取回调日志数据"""
from sqlalchemy import text
try:
# 查询回调日志
result = conn.execute(
result = await db.execute(
text("SELECT request_headers, request_body FROM callback_failure_logs WHERE id = :log_id"),
{"log_id": callback_log_id}
)
@@ -245,8 +251,8 @@ def get_callback_log_data(conn, callback_log_id: int) -> Tuple[bool, Optional[st
return False, None, None
def call_external_api_with_retry(
conn,
async def call_external_api_with_retry(
db: AsyncSession,
request_body: dict,
request_headers: Dict[str, str],
max_retries: int,
@@ -257,16 +263,16 @@ def call_external_api_with_retry(
try:
logger.info(f"🌐 尝试调用外部API接口,第{attempt}次")
with httpx.Client(timeout=30.0) as client:
response = client.post(
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.post(
settings.external_api_url,
json=request_body,
headers=request_headers
)
# 记录推送日志
_log_dtc_push_call(
conn=conn,
await _log_dtc_push_call(
db=db,
callback_failure_log_id=callback_failure_log_id,
request_url=settings.external_api_url,
request_headers=request_headers,
@@ -297,18 +303,20 @@ def call_external_api_with_retry(
# 如果不是最后一次尝试,等待一段时间再重试
if attempt < max_retries:
time.sleep(2 ** attempt) # 指数退避
await asyncio.sleep(2 ** attempt) # 指数退避
# 所有重试都失败了
logger.error(f"❌ 调用外部API接口失败,已重试{max_retries}次")
return False, max_retries
def mark_callback_log_completed(conn, callback_log_id: int) -> bool:
async def mark_callback_log_completed(db: AsyncSession, callback_log_id: int) -> bool:
"""标记CallbackFailureLog记录为已完成"""
from sqlalchemy import text
try:
# 先检查推送日志中是否有响应成功的记录
success_result = conn.execute(
success_result = await db.execute(
text("""
SELECT COUNT(*) as success_count
FROM external_api_logs
@@ -326,7 +334,7 @@ def mark_callback_log_completed(conn, callback_log_id: int) -> bool:
logger.info(f"✅ 回调日志 {callback_log_id} 找到 {success_count} 条成功推送记录,开始标记为完成")
# 更新指定日志记录为已完成
result = conn.execute(
result = await db.execute(
text("""
UPDATE callback_failure_logs
SET is_completed = true
@@ -341,18 +349,18 @@ def mark_callback_log_completed(conn, callback_log_id: int) -> bool:
logger.warning(f"⚠️ 回调日志 {callback_log_id} 已标记为完成或不存在")
return False
conn.commit()
await db.commit()
logger.info(f"✅ 回调日志 {callback_log_id} 已标记为完成")
return True
except Exception as e:
logger.error(f"❌ 标记回调日志 {callback_log_id} 完成失败: {e}", exc_info=True)
conn.rollback()
await db.rollback()
return False
def get_related_records_by_unique_data_list(
conn,
async def get_related_records_by_unique_data_list(
db: AsyncSession,
unique_data_list: list,
callback_log_id: int,
limit_count: int
@@ -360,7 +368,7 @@ def get_related_records_by_unique_data_list(
"""根据unique_data_list中的数据查询相关记录,按创建时间排序
Args:
conn: 数据库连接
db: 异步数据库会话
unique_data_list: 包含手机号、task_id、user_id的数据项列表
callback_log_id: 回调日志ID,作为过滤条件
limit_count: 获取记录数量限制
@@ -370,6 +378,8 @@ def get_related_records_by_unique_data_list(
- 第一个值表示查询是否成功(True表示成功,False表示失败)
- 第二个值为查询到的相关记录列表,失败时返回空列表
"""
from sqlalchemy import text
related_records = []
try:
@@ -387,7 +397,7 @@ def get_related_records_by_unique_data_list(
user_id = item.get('user_id', '')
if phone_number and task_id and user_id:
result = conn.execute(
result = await db.execute(
text("""
SELECT phone_number, task_id, user_id, created_at
FROM callback_failure_data
@@ -427,8 +437,8 @@ def get_related_records_by_unique_data_list(
return False, []
def _log_dtc_push_call(
conn,
async def _log_dtc_push_call(
db: AsyncSession,
callback_failure_log_id: int,
request_url: str,
request_headers: Dict[str, Any],
@@ -439,6 +449,8 @@ def _log_dtc_push_call(
retry_count: int
) -> bool:
"""记录DTC推送调用日志"""
from sqlalchemy import text
try:
# 提取手机号从请求体中
phone_number = None
@@ -448,7 +460,7 @@ def _log_dtc_push_call(
phone_number = first_item['number_data']['number']
# 记录到external_api_logs表
conn.execute(
await db.execute(
text("""
INSERT INTO external_api_logs (
callback_failure_log_id,
@@ -487,11 +499,11 @@ def _log_dtc_push_call(
}
)
conn.commit()
await db.commit()
logger.debug(f"📝 DTC推送调用日志记录成功,状态码: {response_status}")
return True
except Exception as e:
logger.error(f"❌ 记录DTC推送调用日志失败: {e}", exc_info=True)
conn.rollback()
await db.rollback()
return False

View File

@@ -16,7 +16,7 @@ from app.callback_service import (
mark_callback_log_completed,
get_related_records_by_unique_data_list
)
from app.redis_lock import redis_manager
from app.redis_lock import redis_manager, async_redis_manager, distributed_lock
import requests
from app.api_config import API_CONFIG, RETRY_CONFIG
@@ -33,11 +33,18 @@ def conditional_task(enabled=True):
return wrapper
return decorator
# 创建同步数据库连接用于Celery任务
engine = create_engine(
# 创建异步数据库连接用于Celery任务
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
# 使用异步数据库连接
async_engine = create_async_engine(
settings.database_url,
echo=settings.debug,
future=True
future=True,
)
AsyncSessionLocal = async_sessionmaker(
async_engine, class_=AsyncSession, expire_on_commit=False
)
@celery_app.task(bind=True, name='push_data_to_dtc')
@@ -47,27 +54,31 @@ def push_data_to_dtc_task(self):
推送数据给DTC的Celery任务
自动获取一条未完成的回调请求进行处理
"""
import asyncio
from functools import partial
async def async_task():
task_name = 'push_data_to_dtc'
logger.info(f"🌿 开始推送数据给DTC任务")
try:
# 获取分布式锁,使用任务名称作为锁标识
connection_success = redis_manager.connect()
# 连接异步Redis
connection_success = await async_redis_manager.connect()
if not connection_success:
logger.error(f"❌ Redis管理器连接失败,任务停止执行")
raise Exception("Redis连接失败")
lock = redis_manager.create_lock(f"celery_task:{task_name}", timeout=300) # 5分钟超时
# 尝试获取锁
if not lock.acquire(blocking=False):
redis_pool = await async_redis_manager.get_redis_pool()
# 使用异步分布式锁
async with distributed_lock(redis_pool, f"celery_task:{task_name}", timeout=300) as lock_acquired:
if not lock_acquired:
logger.warning(f"⚠️ 任务 {task_name} 正在执行中,跳过本次执行")
return {"status": "skipped", "message": f"任务 {task_name} 正在执行中,跳过本次执行"}
logger.info(f"🔒 成功获取任务 {task_name} 的分布式锁")
with engine.connect() as conn:
logger.info("1")
async with AsyncSessionLocal() as db:
# 更新任务状态
self.update_state(
state='PROGRESS',
@@ -75,10 +86,10 @@ def push_data_to_dtc_task(self):
)
logger.info("2")
# 获取一条未完成的回调请求(按创建时间取最小值)
query_success, callback_log_data = get_uncompleted_callback_log(conn)
query_success, callback_log_data = await get_uncompleted_callback_log(db)
if not query_success:
logger.error("❌ 查询未完成的回调请求失败,尝试再查一次")
query_success, callback_log_data = get_uncompleted_callback_log(conn)
query_success, callback_log_data = await get_uncompleted_callback_log(db)
if not query_success:
logger.error("❌ 再次查询未完成的回调请求失败,停止任务执行")
raise Exception("查询未完成的回调请求失败,任务停止执行")
@@ -103,7 +114,7 @@ def push_data_to_dtc_task(self):
# 保存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)
save_success = await save_callback_data_items(db, data_list, callback_log_id)
if not save_success:
logger.error(f"❌ 保存callback_data:{callback_log_id}失败,停止任务执行")
raise Exception(f"保存callback_data:{callback_log_id}失败,任务停止执行")
@@ -146,8 +157,8 @@ def push_data_to_dtc_task(self):
logger.info(f"🗑️ 移除了 {original_count - len(unique_data_list)} 条重复数据 - callback_log_id: {callback_log_id}, site_id: {site_id}")
# 查询相关数据:根据手机号、task_id、user_id作为条件,按创建时间排序获取前三条记录
query_success, related_records = get_related_records_by_unique_data_list(
conn, unique_data_list, callback_log_id, settings.count_threshold
query_success, related_records = await get_related_records_by_unique_data_list(
db, unique_data_list, callback_log_id, settings.count_threshold
)
if not query_success:
@@ -176,8 +187,8 @@ def push_data_to_dtc_task(self):
meta={'current': 85, 'total': 100, 'status': f'开始推送数据给DTC(ID={callback_log_id}, site_id={site_id})...'}
)
success, retry_count = call_external_api_with_retry(
conn=conn,
success, retry_count = await call_external_api_with_retry(
db=db,
request_body=filtered_request_body,
request_headers=request_headers,
max_retries=settings.external_api_retry_max,
@@ -194,7 +205,7 @@ def push_data_to_dtc_task(self):
logger.info(f"🎉 推送数据给DTC处理完成 - callback_log_id: {callback_log_id}, site_id: {site_id}")
# 标记CallbackFailureLog为已完成
mark_callback_log_completed(conn, callback_log_id)
await mark_callback_log_completed(db, callback_log_id)
# 更新任务状态
self.update_state(
@@ -212,13 +223,25 @@ def push_data_to_dtc_task(self):
return {"status": "error", "message": str(e)}
finally:
# 释放分布式锁
# 分布式锁会通过上下文管理器自动释放
logger.debug(f"🔓 任务 {task_name} 的分布式锁已通过上下文管理器处理")
# 在同步的Celery任务中运行异步代码
try:
if 'lock' in locals():
lock.release()
logger.info(f"🔓 释放任务 {task_name} 的分布式锁")
# 创建新的事件循环
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
return loop.run_until_complete(async_task())
except Exception as e:
logger.error(f"❌ 释放任务 {task_name} 的分布式锁失败: {e}")
logger.error(f"❌ 异步任务执行失败: {e}", exc_info=True)
return {"status": "error", "message": str(e)}
finally:
# 清理事件循环
try:
if 'loop' in locals():
loop.close()
except:
pass
@celery_app.task(bind=True, name='call_api', max_retries=None)

View File

@@ -1,15 +1,212 @@
import redis
import asyncio
import uuid
from typing import Optional
from contextlib import asynccontextmanager
import aioredis
from app.config import settings
from app.logger import get_redis_logger
logger = get_redis_logger()
class AsyncRedisLock:
"""异步Redis分布式锁"""
def __init__(self, redis_pool, key: str, timeout: int = None):
self.redis_pool = redis_pool
self.key = f"lock:{key}"
self.timeout = timeout or settings.redis_lock_timeout
self.identifier = str(uuid.uuid4())
self.acquired = False
async def acquire(self, blocking: bool = False) -> bool:
"""获取分布式锁"""
logger.debug(f"🔒 尝试获取Redis锁: {self.key}")
try:
# 使用SET命令的NX和EX选项原子性地获取锁
result = await self.redis_pool.set(
self.key,
self.identifier,
expire=self.timeout,
exist=self.redis_pool.SET_IF_NOT_EXIST
)
self.acquired = result
if self.acquired:
logger.debug(f"✅ Redis锁获取成功: {self.key}")
else:
logger.debug(f"❌ Redis锁获取失败: {self.key}")
return self.acquired
except Exception as e:
logger.error(f"❌ 获取Redis锁失败: {e}")
return False
async def release(self) -> bool:
"""释放分布式锁"""
if not self.acquired:
return False
try:
# 使用Lua脚本确保只有锁的持有者才能释放锁
lua_script = """
if redis.call("GET", KEYS[1]) == ARGV[1] then
return redis.call("DEL", KEYS[1])
else
return 0
end
"""
result = await self.redis_pool.eval(
lua_script,
1,
self.key,
self.identifier
)
self.acquired = False
released = bool(result)
if released:
logger.debug(f"🔓 Redis锁释放成功: {self.key}")
else:
logger.warning(f"⚠️ Redis锁释放失败,可能已过期: {self.key}")
return released
except Exception as e:
logger.error(f"❌ 释放Redis锁失败: {e}")
return False
async def __aenter__(self):
"""异步上下文管理器入口"""
retries = 0
while retries < settings.redis_lock_max_retries:
if await self.acquire(blocking=False):
return self
await asyncio.sleep(settings.redis_lock_retry_delay)
retries += 1
raise TimeoutError(f"Failed to acquire lock {self.key} after {retries} retries")
async def __aexit__(self, exc_type, exc_val, exc_tb):
"""异步上下文管理器出口"""
await self.release()
@asynccontextmanager
async def distributed_lock(redis_pool, lock_key: str, timeout: int = None):
"""分布式锁上下文管理器
Args:
redis_pool: Redis连接池
lock_key: 锁键名
timeout: 锁超时时间(秒)
Yields:
bool: 是否成功获取锁
"""
full_lock_key = f"lock:{lock_key}"
lock_timeout = timeout or settings.redis_lock_timeout
try:
# 尝试获取锁,使用setnx命令,并设置过期时间
identifier = str(uuid.uuid4())
lock_acquired = await redis_pool.set(
full_lock_key,
identifier,
expire=lock_timeout,
exist=redis_pool.SET_IF_NOT_EXIST
)
logger.debug(f"🔒 尝试获取分布式锁: {full_lock_key}, 结果: {lock_acquired}")
if lock_acquired:
try:
yield True
finally:
# 使用Lua脚本安全释放锁,确保只有锁的持有者才能释放
lua_script = """
if redis.call("GET", KEYS[1]) == ARGV[1] then
return redis.call("DEL", KEYS[1])
else
return 0
end
"""
result = await redis_pool.eval(
lua_script,
1,
full_lock_key,
identifier
)
if result:
logger.debug(f"🔓 分布式锁释放成功: {full_lock_key}")
else:
logger.warning(f"⚠️ 分布式锁释放失败,可能已过期: {full_lock_key}")
else:
yield False
except Exception as e:
logger.error(f"❌ 分布式锁操作失败: {e}")
yield False
class AsyncRedisManager:
def __init__(self):
self.redis_pool: Optional[aioredis.Redis] = None
async def connect(self) -> bool:
"""连接Redis
Returns:
bool: 连接是否成功
"""
try:
logger.info(f"🔴 正在连接Redis: {settings.redis_url}")
self.redis_pool = await aioredis.create_redis_pool(
settings.redis_url,
encoding="utf-8",
minsize=1,
maxsize=settings.redis_max_connections,
timeout=settings.redis_timeout
)
# 测试连接
result = await self.redis_pool.ping()
if result:
logger.info("✅ Redis连接成功")
return True
else:
logger.error("❌ Redis ping失败")
return False
except Exception as e:
logger.error(f"❌ Redis连接失败: {e}")
return False
async def disconnect(self):
"""断开Redis连接"""
if self.redis_pool:
self.redis_pool.close()
await self.redis_pool.wait_closed()
logger.info("🔴 Redis连接已关闭")
async def create_lock(self, key: str, timeout: int = None) -> AsyncRedisLock:
"""创建分布式锁"""
if not self.redis_pool:
raise RuntimeError("Redis client not connected")
return AsyncRedisLock(self.redis_pool, key, timeout)
async def get_redis_pool(self):
"""获取Redis连接池"""
if not self.redis_pool:
await self.connect()
return self.redis_pool
# 为了保持向后兼容,保留同步版本但标记为废弃
class RedisLock:
"""Redis分布式锁"""
def __init__(self, redis_client: redis.Redis, key: str, timeout: int = None):
"""Redis分布式锁 (已废弃,请使用AsyncRedisLock)"""
def __init__(self, redis_client, key: str, timeout: int = None):
self.redis_client = redis_client
self.key = f"lock:{key}"
self.timeout = timeout or settings.redis_lock_timeout
@@ -86,14 +283,15 @@ class RedisLock:
class RedisManager:
def __init__(self):
self.redis_client: Optional[redis.Redis] = None
self.redis_client = None
def connect(self) -> bool:
"""连接Redis
"""连接Redis (已废弃,请使用AsyncRedisManager)
Returns:
bool: 连接是否成功
"""
import redis
try:
logger.info(f"🔴 正在连接Redis: {settings.redis_url}")
self.redis_client = redis.from_url(
@@ -127,5 +325,6 @@ class RedisManager:
return RedisLock(self.redis_client, key, timeout)
# 全局Redis管理器实例
# 全局Redis管理器实例(保持向后兼容)
redis_manager = RedisManager()
async_redis_manager = AsyncRedisManager()

View File

@@ -5,6 +5,7 @@ asyncpg>=0.31.0
alembic>=1.17.2
celery>=5.6.0
redis>=7.1.0
aioredis==1.3.1
flower>=2.0.1
pydantic>=2.12.5
pydantic-settings>=2.12.0
@@ -12,5 +13,4 @@ python-multipart>=0.0.20
httpx>=0.28.1
python-dotenv>=1.2.1
pytest>=9.0.2
pytest-asyncio>=1.3.0
requests==2.32.5