Source code for airflow.providers.amazon.aws.triggers.ecs

# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements.  See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership.  The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License.  You may obtain a copy of the License at
#
#   http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied.  See the License for the
# specific language governing permissions and limitations
# under the License.

from __future__ import annotations

import asyncio
import warnings
from collections.abc import AsyncIterator
from typing import Any

from botocore.exceptions import ClientError, WaiterError

from airflow.exceptions import AirflowProviderDeprecationWarning
from airflow.providers.amazon.aws.hooks.ecs import EcsHook
from airflow.providers.amazon.aws.hooks.logs import AwsLogsHook
from airflow.providers.amazon.aws.triggers.base import AwsBaseWaiterTrigger
from airflow.providers.amazon.aws.utils.task_log_fetcher import AwsTaskLogFetcher, _parse_log_level
from airflow.providers.common.compat.sdk import AirflowException
from airflow.triggers.base import BaseTrigger, TriggerEvent


[docs] class ClusterActiveTrigger(AwsBaseWaiterTrigger): """ Polls the status of a cluster until it's active. :param cluster_arn: ARN of the cluster to watch. :param waiter_delay: The amount of time in seconds to wait between attempts. :param waiter_max_attempts: The number of times to ping for status. Will fail after that many unsuccessful attempts. :param aws_conn_id: The Airflow connection used for AWS credentials. :param region_name: The AWS region where the cluster is located. :param verify: Whether or not to verify SSL certificates. See: https://boto3.amazonaws.com/v1/documentation/api/latest/reference/core/session.html :param botocore_config: Configuration dictionary (key-values) for botocore client. See: https://botocore.amazonaws.com/v1/documentation/api/latest/reference/config.html """
[docs] aws_hook_class = EcsHook
def __init__( self, cluster_arn: str, waiter_delay: int, waiter_max_attempts: int, aws_conn_id: str | None, region_name: str | None = None, verify: bool | str | None = None, botocore_config: dict | None = None, **kwargs, ): super().__init__( serialized_fields={"cluster_arn": cluster_arn}, waiter_name="cluster_active", waiter_args={"clusters": [cluster_arn]}, failure_message="Failure while waiting for cluster to be available", status_message="Cluster is not ready yet", status_queries=["clusters[].status", "failures"], return_key="arn", return_value=cluster_arn, waiter_delay=waiter_delay, waiter_max_attempts=waiter_max_attempts, aws_conn_id=aws_conn_id, region_name=region_name, verify=verify, botocore_config=botocore_config, **kwargs, )
[docs] class ClusterInactiveTrigger(AwsBaseWaiterTrigger): """ Polls the status of a cluster until it's inactive. :param cluster_arn: ARN of the cluster to watch. :param waiter_delay: The amount of time in seconds to wait between attempts. :param waiter_max_attempts: The number of times to ping for status. Will fail after that many unsuccessful attempts. :param aws_conn_id: The Airflow connection used for AWS credentials. :param region_name: The AWS region where the cluster is located. :param verify: Whether or not to verify SSL certificates. See: https://boto3.amazonaws.com/v1/documentation/api/latest/reference/core/session.html :param botocore_config: Configuration dictionary (key-values) for botocore client. See: https://botocore.amazonaws.com/v1/documentation/api/latest/reference/config.html """
[docs] aws_hook_class = EcsHook
def __init__( self, cluster_arn: str, waiter_delay: int, waiter_max_attempts: int, aws_conn_id: str | None, region_name: str | None = None, verify: bool | str | None = None, botocore_config: dict | None = None, **kwargs, ): super().__init__( serialized_fields={"cluster_arn": cluster_arn}, waiter_name="cluster_inactive", waiter_args={"clusters": [cluster_arn]}, failure_message="Failure while waiting for cluster to be deactivated", status_message="Cluster deactivation is not done yet", status_queries=["clusters[].status", "failures"], return_value=cluster_arn, waiter_delay=waiter_delay, waiter_max_attempts=waiter_max_attempts, aws_conn_id=aws_conn_id, region_name=region_name, verify=verify, botocore_config=botocore_config, **kwargs, )
[docs] class TaskDoneTrigger(BaseTrigger): """ Waits for an ECS task to be done, while eventually polling logs. :param cluster: short name or full ARN of the cluster where the task is running. :param task_arn: ARN of the task to watch. :param waiter_delay: The amount of time in seconds to wait between attempts. :param waiter_max_attempts: The number of times to ping for status. Will fail after that many unsuccessful attempts. :param aws_conn_id: The Airflow connection used for AWS credentials. :param region_name: The AWS region where the cluster is located. Used to build the hook. :param log_region_name: The AWS region where the CloudWatch logs are stored. Defaults to ``region_name`` when not set, which is correct unless the task definition ships its logs to another region. :param verify: Whether or not to verify SSL certificates. Used to build the hook. :param botocore_config: Configuration dictionary for the botocore client. Used to build the hook. :param region: (deprecated) use ``region_name`` instead. """ def __init__( self, cluster: str, task_arn: str, waiter_delay: int, waiter_max_attempts: int, aws_conn_id: str | None, region_name: str | None = None, log_group: str | None = None, log_stream: str | None = None, verify: bool | str | None = None, botocore_config: dict | None = None, region: str | None = None, log_region_name: str | None = None, ): if region is not None: warnings.warn( "`region` is deprecated and will be removed in a future release. " "Please use `region_name` instead.", AirflowProviderDeprecationWarning, stacklevel=2, ) region_name = region
[docs] self.cluster = cluster
[docs] self.task_arn = task_arn
[docs] self.waiter_delay = waiter_delay
[docs] self.waiter_max_attempts = waiter_max_attempts
[docs] self.aws_conn_id = aws_conn_id
[docs] self.region_name = region_name
[docs] self.verify = verify
[docs] self.botocore_config = botocore_config
[docs] self.log_group = log_group
[docs] self.log_stream = log_stream
[docs] self.log_region_name = log_region_name
[docs] def serialize(self) -> tuple[str, dict[str, Any]]: return ( self.__class__.__module__ + "." + self.__class__.__qualname__, { "cluster": self.cluster, "task_arn": self.task_arn, "waiter_delay": self.waiter_delay, "waiter_max_attempts": self.waiter_max_attempts, "aws_conn_id": self.aws_conn_id, "region_name": self.region_name, "log_region_name": self.log_region_name, "log_group": self.log_group, "log_stream": self.log_stream, "verify": self.verify, "botocore_config": self.botocore_config, }, )
[docs] async def run(self) -> AsyncIterator[TriggerEvent]: # Triggers serialized before ``log_region_name`` existed deserialize without it, so an # unset value keeps reading logs from the cluster region as before. log_region_name = self.log_region_name if self.log_region_name is not None else self.region_name async with ( await EcsHook( aws_conn_id=self.aws_conn_id, region_name=self.region_name, verify=self.verify, config=self.botocore_config, ).get_async_conn() as ecs_client, await AwsLogsHook( aws_conn_id=self.aws_conn_id, region_name=log_region_name, verify=self.verify, config=self.botocore_config, ).get_async_conn() as logs_client, ): waiter = ecs_client.get_waiter("tasks_stopped") logs_token = None while self.waiter_max_attempts: self.waiter_max_attempts -= 1 try: await waiter.wait( cluster=self.cluster, tasks=[self.task_arn], WaiterConfig={"MaxAttempts": 1} ) # we reach this point only if the waiter met a success criteria yield TriggerEvent( {"status": "success", "task_arn": self.task_arn, "cluster": self.cluster} ) return except WaiterError as error: if "terminal failure" in str(error): raise self.log.info("Status of the task is %s", error.last_response["tasks"][0]["lastStatus"]) await asyncio.sleep(int(self.waiter_delay)) finally: if self.log_group and self.log_stream: logs_token = await self._forward_logs(logs_client, logs_token) raise AirflowException("Waiter error: max attempts reached")
async def _forward_logs(self, logs_client, next_token: str | None = None) -> str | None: """ Read logs from the cloudwatch stream and print them to the task logs. :return: the token to pass to the next iteration to resume where we started. """ while True: if next_token is not None: token_arg: dict[str, str] = {"nextToken": next_token} else: token_arg = {} try: response = await logs_client.get_log_events( logGroupName=self.log_group, logStreamName=self.log_stream, startFromHead=True, **token_arg, ) except ClientError as ce: if ce.response["Error"]["Code"] == "ResourceNotFoundException": self.log.info( "Tried to get logs from stream %s in group %s but it didn't exist (yet). " "Will try again.", self.log_stream, self.log_group, ) return None raise events = response["events"] for log_event in events: level = _parse_log_level(log_event["message"]) self.log.log(level, AwsTaskLogFetcher.event_to_str(log_event)) if len(events) == 0 or next_token == response["nextForwardToken"]: return response["nextForwardToken"] next_token = response["nextForwardToken"]

Was this entry helpful?