如何在Django应用中用pytest测试Celery Chord?遇数据库连接问题
Celery Chord测试问题及解决方案
依赖版本
celery==5.2.7 django-celery-results==2.4.0 django==4.1 pytest==7.1.2 pytest-django==4.5.2 pytest-celery==0.0.0
问题场景
要测试名为start_task的Celery任务,该任务会创建包含N个work_task的chord,并通过summarize_task作为回调汇总结果。
初始测试问题
使用仅包含db fixture的测试代码:
def test_function(db): ... obj = make_obj() ... start_task.delay(obj)
- 单个
work_task可成功执行,但chord始终无法完成,summarize_task完全不触发。
添加Celery Fixture后的问题
修改测试代码引入celery_app和celery_worker fixture:
def test_function(db, celery_app, celery_worker): ... obj = make_obj() ... start_task.delay(obj)
- 执行
make_obj时直接失败,报错:E psycopg2.InterfaceError: connection already closed
目前临时方案是手动调用任务绕开Celery,但无法验证chord机制,仅能测试任务逻辑。
解决方案
1. 更换合适的Broker
默认内存Broker(memory://)对chord支持有限,建议切换到Redis或RabbitMQ,在Django配置文件中修改:
CELERY_BROKER_URL = 'redis://localhost:6379/0' # 或RabbitMQ地址 CELERY_RESULT_BACKEND = 'django-db'
2. 修复数据库连接泄漏问题
- 配置Celery自动回收连接:
在Django配置中添加:CELERY_WORKER_MAX_TASKS_PER_CHILD = 1 # 每个worker进程处理完任务就重启,避免连接失效 CELERY_DB_REUSE_MAX = None # 禁用连接复用,每次任务重新建立连接 - 任务内手动管理连接:
在Celery任务中显式处理数据库连接:from django.db import connection from celery import shared_task @shared_task def work_task(obj_id): try: obj = YourModel.objects.get(id=obj_id) # 执行任务逻辑 ... finally: if not connection.is_usable(): connection.close() - 测试中使用正确的数据库标记:
给测试函数添加@pytest.mark.django_db(transaction=True)标记,确保测试事务正确管理:import pytest @pytest.mark.django_db(transaction=True) def test_function(db, celery_app, celery_worker): ...
3. 正确等待Chord执行完成
测试时不要仅调用delay(),需主动等待任务完成并验证结果:
def test_chord_execution(db, celery_app, celery_worker): obj = make_obj() # 建议传对象ID而非对象本身,避免序列化问题 task_result = start_task.delay(obj.id) # 等待chord执行完成,设置合理超时时间 final_summary = task_result.get(timeout=15) # 断言汇总结果符合预期 assert final_summary == expected_value
4. 直接测试Chord构造逻辑
如果start_task仅负责构造chord,可以直接在测试中手动构建chord,验证其执行流程:
from celery import chord def test_direct_chord(db, celery_app, celery_worker): obj = make_obj() # 构建work_task任务组,模拟N个任务 task_group = [work_task.s(obj.id) for _ in range(5)] # 创建chord并执行 chord_result = chord(task_group)(summarize_task.s()) # 等待回调执行完成 summary = chord_result.get(timeout=10) # 验证结果 assert len(summary) == 5
内容的提问来源于stack exchange,提问作者boatcoder
相关产品推荐
相关产品推荐

