You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Celery任务ECS终止问题:需更新装饰器处理ProtectionEnabled状态变更

问题描述

我有一个基于Django的应用,在AWS Elastic Container Service(ECS)上运行多个Celery任务,用SQS做消息中间件。当前遇到的问题:前一个任务完成后,Celery会在现有ECS任务里启动新任务,但我的装饰器会把ProtectionEnabled状态从true改成false,导致20秒后ECS任务被终止,新任务无法运行。

我通过CloudWatch告警监控消息队列,用来终止已完成的ECS任务。现在想更新装饰器,让新Celery任务启动时把ProtectionEnabled从false改回true,但不知道怎么实现。

启动Celery任务的命令:

celery -A myapp_settings.celery worker --concurrency=1 l info -Q sqs-celery

相关代码

container_decorator.py

class ContainerAgent:
    class Error(Exception):
        pass

    class RequestError(Error, IOError):
        pass

    def __init__(
        self,
        ecs_agent_uri: str,
        timeout: int = 10,
        session: requests.Session = None,
        logger: logging.Logger = None,
    ) -> None:
        self._ecs_agent_uri = ecs_agent_uri
        self._timeout = timeout

        self._session = session or requests.Session()
        self._logger = logger or logging.getLogger(self.__class__.__name__)

    def _request(self, *, path: str, data: Optional[dict] = None) -> dict:
        url = f"{self._ecs_agent_uri}{path}"
        self._logger.info(f"Performing request... {url=} {data=}")

        try:
            response = self._session.put(url=url, json=data, timeout=self._timeout)
            self._logger.info(f"Got response. {response.status_code=} {response.content=}")

            response.raise_for_status()
            return response.json()
        except requests.RequestException as e:
            response_body = e.response.text if e.response is not None else None
            self._logger.warning(f"Request error! {url=} {data=} {e=} {response_body=}")

            raise self.RequestError(str(e)) from e

    def toggle_scale_in_protection(self, *, enable: bool = True, expire_in_minutes: int = 2880):
        response = self._request(
            path="/task-protection/v1/state",
            data={"ProtectionEnabled": enable, "ExpiresInMinutes": expire_in_minutes},
        )

        try:
            return response["protection"]["ProtectionEnabled"]
        except KeyError as e:
            raise self.Error(f"Task scale-in protection endpoint error: {response=}") from e


def enable_scale_in_protection(*, logger: logging.Logger = None):
    def decorator(f):
        if not (ecs_agent_uri := os.getenv("ECS_AGENT_URI")):
            (logger or logging).warning(f"Scale-in protection not enabled. {ecs_agent_uri=}")
            return f

        client = ContainerAgent(ecs_agent_uri=ecs_agent_uri, logger=logger)

        @wraps(f)
        def wrapper(*args, **kwargs):
            try:
                client.toggle_scale_in_protection(enable=True)
            except client.Error as e:
                (logger or logging).warning(f"Scale-in protection not enabled. {e}")
                protection_set = False
            else:
                protection_set = True

            try:
                return f(*args, **kwargs)
            finally:
                if protection_set:
                    client.toggle_scale_in_protection(enable=False)

        return wrapper
    return decorator

celery_tasks.py

@shared_task(name="add_spider_schedule", base=AbortableTask)
@enable_scale_in_protection(logger=get_task_logger(__name__))
def add_spider_schedule(user_id, spider_id):
    settings_module = os.environ.get('DJANGO_SETTINGS_MODULE')
    if settings_module == 'myapp_settings.settings.production':
        return add_spider_schedule_production(user_id, spider_id, add_spider_schedule)
    else:
        return print('Unknown settings module')

def add_spider_schedule_production(user_id, spider_id, task_object):
    """
    Adds the schedule for the specified spider.
    :param spider_id: The ID of the spider to schedule.
    :return: A string representation of the spider and task IDs.
    """
    # 日志配置:将所有print输出重定向到任务日志
    logger = get_task_logger(task_object.request.id)
    old_outs = sys.stdout, sys.stderr
    rlevel = add_spider_schedule.app.conf.worker_redirect_stdouts_level
    add_spider_schedule.app.log.redirect_stdouts_to_logger(logger, rlevel)

    # 获取Spider模型实例
    spider = Spider.objects.get(id=spider_id)
    # 获取当前用户
    user = User.objects.get(id=user_id)

    # 从模型中获取相关文件名称
    spider_config_file = spider.spider_config_file.file
    yaml_config_file = spider.yaml_config_file.file
    template_file = spider.template_file.file
    mongodb_database_name = spider.mongodb_collection.database_name
    mongodb_collection_name = spider.mongodb_collection.collection_name

    # 从S3存储桶读取文件内容
    spider_config_file_contents = load_content_from_s3(AWS_STORAGE_BUCKET_NAME, rf"{PUBLIC_MEDIA_LOCATION}/{spider_config_file}")
    yaml_config_path = load_content_from_s3(AWS_STORAGE_BUCKET_NAME, rf"{PUBLIC_MEDIA_LOCATION}/{yaml_config_file}")
    input_file_path = load_content_from_s3(AWS_STORAGE_BUCKET_NAME, rf"{PUBLIC_MEDIA_LOCATION}/{template_file}")

    # 将JSON格式的关键字参数转换为字典
    kwargs = json.loads(spider.kwargs) if spider.kwargs else {}

    # 从spider_config_file内容创建模块
    spider_module = import_module(spider_config_file_contents, "spider_config")
 
    is_scraping_finished = False

    async def run_spider():
        try:
            await spider_module.run(
                yaml_config_path=yaml_config_path,
                input_file_path=input_file_path,
                mongodb_name=mongodb_database_name,
                mongodb_collection_name=mongodb_collection_name,
                task_object=task_object,
                mode="sf-lab",
                **kwargs
            )
            nonlocal is_scraping_finished
            is_scraping_finished = True
        except Exception as e:
            raise Exception(f"运行爬虫时出错: {e}")

    async def check_if_aborted():
        while True:
            if task_object.is_aborted():
                print("检测到任务已取消")
                raise Exception("任务已取消")
            elif is_scraping_finished:
                break
            await asyncio.sleep(0.1)

    loop = asyncio.get_event_loop()
    loop.run_until_complete(asyncio.gather(run_spider(), check_if_aborted()))

    sys.stdout, sys.stderr = old_outs  # 恢复标准输出

    return f"[spider: {spider_id}, task_id: {task_object.request.id}]"
解决方案

问题核心是原装饰器在单个任务结束后强制关闭了ECS任务的保护,但Celery Worker会复用ECS任务处理下一个任务,导致新任务启动前ECS任务已处于无保护状态,被CloudWatch告警终止。

需要调整装饰器逻辑,放弃"任务结束后立即关闭保护"的做法,改为在每个任务启动前自动开启保护,让保护自动过期而非主动关闭。修改后的代码如下:

from celery.signals import task_prerun, task_postrun
from functools import wraps
import os
import logging
from typing import Optional
import requests

class ContainerAgent:
    # 保持原有ContainerAgent代码不变
    class Error(Exception):
        pass

    class RequestError(Error, IOError):
        pass

    def __init__(
        self,
        ecs_agent_uri: str,
        timeout: int = 10,
        session: requests.Session = None,
        logger: logging.Logger = None,
    ) -> None:
        self._ecs_agent_uri = ecs_agent_uri
        self._timeout = timeout

        self._session = session or requests.Session()
        self._logger = logger or logging.getLogger(self.__class__.__name__)

    def _request(self, *, path: str, data: Optional[dict] = None) -> dict:
        url = f"{self._ecs_agent_uri}{path}"
        self._logger.info(f"执行请求... {url=} {data=}")

        try:
            response = self._session.put(url=url, json=data, timeout=self._timeout)
            self._logger.info(f"收到响应. {response.status_code=} {response.content=}")

            response.raise_for_status()
            return response.json()
        except requests.RequestException as e:
            response_body = e.response.text if e.response is not None else None
            self._logger.warning(f"请求错误! {url=} {data=} {e=} {response_body=}")

            raise self.RequestError(str(e)) from e

    def toggle_scale_in_protection(self, *, enable: bool = True, expire_in_minutes: int = 2880):
        response = self._request(
            path="/task-protection/v1/state",
            data={"ProtectionEnabled": enable, "ExpiresInMinutes": expire_in_minutes},
        )

        try:
            return response["protection"]["ProtectionEnabled"]
        except KeyError as e:
            raise self.Error(f"任务缩容保护端点错误: {response=}") from e


def enable_scale_in_protection(*, logger: logging.Logger = None):
    def decorator(f):
        if not (ecs_agent_uri := os.getenv("ECS_AGENT_URI")):
            (logger or logging).warning(f"未启用缩容保护. {ecs_agent_uri=}")
            return f

        client = ContainerAgent(ecs_agent_uri=ecs_agent_uri, logger=logger)
        protection_expire_minutes = 2880  # 保护默认有效期48小时,可根据需求调整

        # 任务启动前自动开启缩容保护
        @task_prerun.connect(sender=f)
        def on_task_prerun(sender=None, task_id=None, **kwargs):
            try:
                client.toggle_scale_in_protection(enable=True, expire_in_minutes=protection_expire_minutes)
                logger.info(f"为任务 {task_id} 启用缩容保护")
            except client.Error as e:
                logger.warning(f"为任务 {task_id} 启用缩容保护失败: {e}")

        # 任务结束后不主动关闭保护,依赖自动过期机制
        # 可选:如果希望Worker空闲时也短暂保持保护,可在此处添加续期逻辑
        @task_postrun.connect(sender=f)
        def on_task_postrun(sender=None, task_id=None, **kwargs):
            # try:
            #     client.toggle_scale_in_protection(enable=True, expire_in_minutes=10)  # 续期10分钟
            #     logger.info(f"任务 {task_id} 结束后,延长Worker缩容保护")
            # except client.Error as e:
            #     logger.warning(f"延长缩容保护失败: {e}")
            pass

        @wraps(f)
        def wrapper(*args, **kwargs):
            return f(*args, **kwargs)

        return wrapper
    return decorator

关键改动说明

  1. 利用Celery的task_prerun信号,在每个任务启动前自动开启ECS任务的缩容保护,确保新任务启动时ECS任务处于受保护状态。
  2. 移除原装饰器中任务结束后强制关闭保护的逻辑,改为让保护自动过期,避免Worker复用ECS任务时出现保护真空期。
  3. 可根据业务需求调整保护的过期时间:如果任务排队频繁,可设置较长有效期;如果Worker空闲时间久,可缩短有效期,配合CloudWatch告警及时终止空闲ECS任务。

额外优化建议

  • 调整CloudWatch告警触发条件,仅当队列无待处理任务且ECS任务无缩容保护时才终止任务,避免误杀正在处理任务的ECS实例。
  • 可添加worker_ready信号处理函数,在Worker启动时就开启一次保护,防止Worker刚启动还未接收任务就被终止。

内容的提问来源于stack exchange,提问作者Adrian

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.05 19:07:33