# 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
from collections.abc import AsyncIterator
from typing import Any
from airflow.exceptions import AirflowException
from airflow.providers.amazon.aws.hooks.base_aws import AwsGenericHook
from airflow.providers.amazon.aws.utils.waiter_with_logging import async_wait
from airflow.triggers.base import BaseTrigger, TriggerEvent
from airflow.utils.helpers import prune_dict
[docs]
class AwsBaseWaiterTrigger(BaseTrigger):
"""
Base class for all AWS Triggers that follow the "standard" model of just waiting on a waiter.
Subclasses should set the ``aws_hook_class`` attribute to the hook they need. The hook is then
built from the parameters this class already serializes, so the deferred half of a task talks to
AWS with the same region, SSL verification setting and botocore configuration as the synchronous
half. Subclasses whose hook takes something else may override :meth:`_hook_parameters` or, as a
last resort, :meth:`hook` itself.
:param serialized_fields: Fields that are specific to the subclass trigger and need to be serialized
to be passed to the __init__ method on deserialization.
The conn id, region, and waiter delay & attempts are always serialized.
format: {<parameter_name>: <parameter_value>}
:param waiter_name: The name of the (possibly custom) boto waiter to use.
:param waiter_args: The arguments to pass to the waiter.
:param failure_message: The message to log if a failure state is reached.
:param status_message: The message logged when printing the status of the service.
:param status_queries: A list containing the JMESPath queries to retrieve status information from
the waiter response. See https://jmespath.org/tutorial.html
:param return_key: The key to use for the return_value in the TriggerEvent this emits on success.
Defaults to "value".
:param return_value: A value that'll be returned in the return_key field of the TriggerEvent.
Set to None if there is nothing to return.
:param waiter_delay: The amount of time in seconds to wait between attempts.
:param waiter_max_attempts: The maximum number of attempts to be made.
:param waiter_config_overrides: A dict to update waiter's default configuration. Only specified keys will
be updated.
:param aws_conn_id: The Airflow connection used for AWS credentials. To be used to build the hook.
:param region_name: The AWS region where the resources to watch are. To be used to build the hook.
:param verify: Whether or not to verify SSL certificates. To be used to build the hook.
See: https://boto3.amazonaws.com/v1/documentation/api/latest/reference/core/session.html
:param botocore_config: Configuration dictionary (key-values) for botocore client.
To be used to build the hook. For available key-values see:
https://botocore.amazonaws.com/v1/documentation/api/latest/reference/config.html
"""
# Should be assigned in child class, unless hook() is overridden.
[docs]
aws_hook_class: type[AwsGenericHook]
[docs]
def __init_subclass__(cls, **kwargs: Any) -> None:
"""
Reject a subclass that cannot build a hook, at import time.
``hook()`` is only reached from ``run()``, which executes in the triggerer, so without this
a subclass that declares neither would defer successfully and fail later, out of sight of
the task that deferred. The operator side gets the same guarantee from
``AwsBaseHookMixin.validate_attributes``. This runs on class creation rather than in
``__init__`` because subclasses such as ``EksDeleteClusterTrigger`` never call
``super().__init__()``.
"""
super().__init_subclass__(**kwargs)
if cls.hook is not AwsBaseWaiterTrigger.hook:
return
hook_class = getattr(cls, "aws_hook_class", None)
if hook_class is None:
raise AttributeError(
f"Class attribute '{cls.__name__}.aws_hook_class' should be set, "
f"or {cls.__name__}.hook() overridden."
)
if not (isinstance(hook_class, type) and issubclass(hook_class, AwsGenericHook)):
raise AttributeError(
f"Class attribute '{cls.__name__}.aws_hook_class' is not a subclass of AwsGenericHook."
)
def __init__(
self,
*,
serialized_fields: dict[str, Any],
waiter_name: str,
waiter_args: dict[str, Any],
failure_message: str,
status_message: str,
status_queries: list[str],
return_key: str = "value",
return_value: Any,
waiter_delay: int,
waiter_max_attempts: int,
waiter_config_overrides: dict[str, Any] | None = None,
aws_conn_id: str | None = "aws_default",
region_name: str | None = None,
verify: bool | str | None = None,
botocore_config: dict | None = None,
):
super().__init__()
# parameters that should be hardcoded in the child's implem
[docs]
self.serialized_fields = serialized_fields
[docs]
self.waiter_name = waiter_name
[docs]
self.waiter_args = waiter_args
[docs]
self.failure_message = failure_message
[docs]
self.status_message = status_message
[docs]
self.status_queries = status_queries
[docs]
self.waiter_config_overrides = waiter_config_overrides
[docs]
self.return_key = return_key
[docs]
self.return_value = return_value
# parameters that should be passed directly from the child's parameters
[docs]
self.waiter_delay = waiter_delay
[docs]
self.attempts = waiter_max_attempts
[docs]
self.aws_conn_id = aws_conn_id
[docs]
self.region_name = region_name
[docs]
self.botocore_config = botocore_config
[docs]
def serialize(self) -> tuple[str, dict[str, Any]]:
# here we put together the "common" params,
# and whatever extras we got from the subclass in serialized_fields
params = dict(
{
"waiter_delay": self.waiter_delay,
"waiter_max_attempts": self.attempts,
"aws_conn_id": self.aws_conn_id,
},
**self.serialized_fields,
)
# if we serialize the None value from this, it breaks subclasses that don't have it in their ctor.
params.update(
prune_dict(
{
# Keep previous behaviour when empty string in region_name evaluated as `None`
"region_name": self.region_name or None,
"verify": self.verify,
"botocore_config": self.botocore_config,
}
)
)
return (
# remember that self is an instance of the subclass here, not of this class.
self.__class__.__module__ + "." + self.__class__.__qualname__,
params,
)
@property
def _hook_parameters(self) -> dict[str, Any]:
"""Mapping of the serialized parameters onto the hook's constructor keywords."""
return {
"aws_conn_id": self.aws_conn_id,
"region_name": self.region_name,
"verify": self.verify,
"config": self.botocore_config,
}
[docs]
def hook(self) -> AwsGenericHook:
"""Build the hook this trigger waits with."""
if not hasattr(self, "aws_hook_class"):
raise AttributeError(
f"Class attribute '{type(self).__name__}.aws_hook_class' should be set, "
f"or {type(self).__name__}.hook() overridden."
)
return self.aws_hook_class(**self._hook_parameters)
def _event_from_exception(self, error: AirflowException) -> TriggerEvent:
return TriggerEvent(
{
"status": "error",
"message": str(error),
self.return_key: self.return_value,
}
)
[docs]
async def run(self) -> AsyncIterator[TriggerEvent]:
hook = self.hook()
async with await hook.get_async_conn() as client:
waiter = hook.get_waiter(
self.waiter_name,
deferrable=True,
client=client,
config_overrides=self.waiter_config_overrides,
)
try:
await async_wait(
waiter,
self.waiter_delay,
self.attempts,
self.waiter_args,
self.failure_message,
self.status_message,
self.status_queries,
)
except AirflowException as error:
yield self._event_from_exception(error)
else:
yield TriggerEvent({"status": "success", self.return_key: self.return_value})