This is an automated email from the ASF dual-hosted git repository.
o-nikolas pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new a80151b1683 Add deferrable mode to EmrServerlessJobSensor (#71652)
a80151b1683 is described below
commit a80151b1683a45a3db36a3736f7476cf09f2a2af
Author: saadbelgi <[email protected]>
AuthorDate: Wed Sep 16 06:18:47 2026 +0530
Add deferrable mode to EmrServerlessJobSensor (#71652)
Waiting on a long-running EMR Serverless job held a worker slot for the
entire duration of the job.. The dedicated trigger overrides the shared
waiter's acceptors rather than adding waiter JSON per target-state
combination, because the sensor's target_states are only known at runtime
and failure states must keep taking precedence over them the way
poke mode already does.
---
.../amazon/docs/operators/emr/emr_serverless.rst | 2 +
.../airflow/providers/amazon/aws/sensors/emr.py | 35 ++++++++-
.../airflow/providers/amazon/aws/triggers/emr.py | 76 ++++++++++++++++++++
.../amazon/aws/sensors/test_emr_serverless_job.py | 54 +++++++++++++-
.../tests/unit/amazon/aws/triggers/test_emr.py | 82 +++++++++++++++++++++-
.../unit/amazon/aws/triggers/test_serialization.py | 9 +++
6 files changed, 254 insertions(+), 4 deletions(-)
diff --git a/providers/amazon/docs/operators/emr/emr_serverless.rst
b/providers/amazon/docs/operators/emr/emr_serverless.rst
index 761713c1b08..004c825f8e2 100644
--- a/providers/amazon/docs/operators/emr/emr_serverless.rst
+++ b/providers/amazon/docs/operators/emr/emr_serverless.rst
@@ -130,6 +130,8 @@ Wait on an EMR Serverless Job state
To monitor the state of an EMR Serverless Job you can use
:class:`~airflow.providers.amazon.aws.sensors.emr.EmrServerlessJobSensor`.
+This sensor can be run in deferrable mode by passing ``deferrable=True`` as a
parameter. This requires
+the aiobotocore module to be installed.
.. exampleinclude::
/../../amazon/tests/system/amazon/aws/example_emr_serverless.py
:language: python
diff --git a/providers/amazon/src/airflow/providers/amazon/aws/sensors/emr.py
b/providers/amazon/src/airflow/providers/amazon/aws/sensors/emr.py
index cfb8575a752..8e2536e7881 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/sensors/emr.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/sensors/emr.py
@@ -26,6 +26,7 @@ from airflow.providers.amazon.aws.links.emr import
EmrClusterLink, EmrLogsLink,
from airflow.providers.amazon.aws.sensors.base_aws import AwsBaseSensor
from airflow.providers.amazon.aws.triggers.emr import (
EmrContainerTrigger,
+ EmrServerlessJobSensorTrigger,
EmrStepSensorTrigger,
EmrTerminateJobFlowTrigger,
)
@@ -115,7 +116,7 @@ class EmrBaseSensor(AwsBaseSensor[EmrHook]):
class EmrServerlessJobSensor(AwsBaseSensor[EmrServerlessHook]):
"""
- Poll the state of the job run until it reaches a terminal state; fails if
the job run fails.
+ Poll the state of the job run until it reaches one of the target states;
fails if the job run fails.
.. seealso::
For more information on how to use this sensor, take a look at the
guide:
@@ -123,7 +124,8 @@ class
EmrServerlessJobSensor(AwsBaseSensor[EmrServerlessHook]):
:param application_id: application_id to check the state of
:param job_run_id: job_run_id to check the state of
- :param target_states: a set of states to wait for, defaults to 'SUCCESS'
+ :param target_states: a set of states to wait for, defaults to ``SUCCESS``.
+ :param deferrable: Run sensor in the deferrable mode.
:param aws_conn_id: The Airflow connection used for AWS credentials.
If this is ``None`` or empty then the default boto3 behaviour is used.
If
running Airflow in a distributed manner and aws_conn_id is None or
@@ -146,13 +148,42 @@ class
EmrServerlessJobSensor(AwsBaseSensor[EmrServerlessHook]):
application_id: str,
job_run_id: str,
target_states: set | frozenset =
frozenset(EmrServerlessHook.JOB_SUCCESS_STATES),
+ deferrable: bool = conf.getboolean("operators", "default_deferrable",
fallback=False),
**kwargs: Any,
) -> None:
self.target_states = target_states
self.application_id = application_id
self.job_run_id = job_run_id
+ self.deferrable = deferrable
super().__init__(**kwargs)
+ def execute(self, context: Context) -> None:
+ if not self.deferrable:
+ super().execute(context=context)
+ elif not self.poke(context):
+ self.defer(
+ timeout=timedelta(seconds=self.timeout),
+ trigger=EmrServerlessJobSensorTrigger(
+ application_id=self.application_id,
+ job_run_id=self.job_run_id,
+ target_states=self.target_states,
+ waiter_delay=int(self.poke_interval),
+ aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ botocore_config=self.botocore_config,
+ ),
+ method_name="execute_complete",
+ )
+
+ def execute_complete(self, context: Context, event: dict[str, Any] | None
= None) -> None:
+ validated_event = validate_execute_complete_event(event)
+
+ if validated_event["status"] != "success":
+ raise RuntimeError(f"Error while running job: {validated_event}")
+
+ self.log.info("EMR Serverless job %s reached a target state.",
self.job_run_id)
+
def poke(self, context: Context) -> bool:
response =
self.hook.conn.get_job_run(applicationId=self.application_id,
jobRunId=self.job_run_id)
diff --git a/providers/amazon/src/airflow/providers/amazon/aws/triggers/emr.py
b/providers/amazon/src/airflow/providers/amazon/aws/triggers/emr.py
index 48ddb4b0c31..bd545affd57 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/triggers/emr.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/triggers/emr.py
@@ -371,6 +371,82 @@ class
EmrServerlessStopApplicationTrigger(AwsBaseWaiterTrigger):
return EmrServerlessHook(self.aws_conn_id)
+class EmrServerlessJobSensorTrigger(AwsBaseWaiterTrigger):
+ """
+ Poll an EMR Serverless job run until it reaches a target or failure state.
+
+ :param application_id: The ID of the application the job is running on.
+ :param job_run_id: The ID of the job run.
+ :param target_states: The states that indicate the sensor has succeeded.
+ :param waiter_delay: The time in seconds to wait between polling attempts.
+ :param waiter_max_attempts: The maximum number of attempts to be made.
Defaults to an infinite wait.
+ :param aws_conn_id: Reference to the AWS connection ID.
+ :param region_name: The AWS region where the job is running.
+ :param verify: Whether to verify SSL certificates.
+ :param botocore_config: Configuration dictionary for the botocore client.
+ """
+
+ def __init__(
+ self,
+ application_id: str,
+ job_run_id: str,
+ target_states: set[str] | frozenset[str],
+ waiter_delay: int = 60,
+ waiter_max_attempts: int = sys.maxsize,
+ aws_conn_id: str | None = "aws_default",
+ region_name: str | None = None,
+ verify: bool | str | None = None,
+ botocore_config: dict | None = None,
+ ) -> None:
+ normalized_target_states = set(target_states)
+ super().__init__(
+ serialized_fields={
+ "application_id": application_id,
+ "job_run_id": job_run_id,
+ "target_states": normalized_target_states,
+ },
+ waiter_name="serverless_job_completed",
+ waiter_args={"applicationId": application_id, "jobRunId":
job_run_id},
+ failure_message="EMR Serverless job failed",
+ status_message="EMR Serverless job status is",
+ status_queries=["jobRun.state", "jobRun.stateDetails"],
+ return_value=None,
+ waiter_delay=waiter_delay,
+ waiter_max_attempts=waiter_max_attempts,
+ waiter_config_overrides={"acceptors":
self._build_waiter_acceptors(normalized_target_states)},
+ aws_conn_id=aws_conn_id,
+ region_name=region_name,
+ verify=verify,
+ botocore_config=botocore_config,
+ )
+
+ @staticmethod
+ def _build_waiter_acceptors(target_states: set[str]) -> list[dict[str,
str]]:
+ acceptors = []
+ for states, waiter_state in (
+ (EmrServerlessHook.JOB_FAILURE_STATES, "failure"),
+ (target_states, "success"),
+ ):
+ for state in states:
+ acceptors.append(
+ {
+ "matcher": "path",
+ "argument": "jobRun.state",
+ "expected": state,
+ "state": waiter_state,
+ }
+ )
+ return acceptors
+
+ def hook(self) -> EmrServerlessHook:
+ return EmrServerlessHook(
+ aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ config=self.botocore_config,
+ )
+
+
class EmrServerlessStartJobTrigger(AwsBaseWaiterTrigger):
"""
Poll an Emr Serverless job run and wait for it to be completed.
diff --git
a/providers/amazon/tests/unit/amazon/aws/sensors/test_emr_serverless_job.py
b/providers/amazon/tests/unit/amazon/aws/sensors/test_emr_serverless_job.py
index 47e92082a17..072e5c72242 100644
--- a/providers/amazon/tests/unit/amazon/aws/sensors/test_emr_serverless_job.py
+++ b/providers/amazon/tests/unit/amazon/aws/sensors/test_emr_serverless_job.py
@@ -17,12 +17,15 @@
# under the License.
from __future__ import annotations
+from datetime import timedelta
+from unittest import mock
from unittest.mock import MagicMock
import pytest
from airflow.providers.amazon.aws.sensors.emr import EmrServerlessJobSensor
-from airflow.providers.common.compat.sdk import AirflowException
+from airflow.providers.amazon.aws.triggers.emr import
EmrServerlessJobSensorTrigger
+from airflow.providers.common.compat.sdk import AirflowException, TaskDeferred
class TestEmrServerlessJobSensor:
@@ -78,3 +81,52 @@ class
TestPokeRaisesAirflowException(TestEmrServerlessJobSensor):
assert exception_msg == str(ctx.value)
self.assert_get_job_run_was_called_once_with_app_and_run_id()
+
+
+class TestEmrServerlessJobSensorDeferrable(TestEmrServerlessJobSensor):
+ def test_sensor_defer_trigger_parameters(self):
+ sensor = EmrServerlessJobSensor(
+ task_id="test_emr_serverless_job_sensor",
+ application_id=self.app_id,
+ job_run_id=self.job_run_id,
+ target_states={"RUNNING"},
+ aws_conn_id="aws_default",
+ region_name="eu-west-1",
+ verify=False,
+ botocore_config={"read_timeout": 42},
+ deferrable=True,
+ poke_interval=10,
+ timeout=300,
+ )
+
+ with mock.patch.object(EmrServerlessJobSensor, "poke", autospec=True,
return_value=False):
+ with pytest.raises(TaskDeferred) as exc:
+ sensor.execute(context=None)
+
+ trigger = exc.value.trigger
+ assert isinstance(trigger, EmrServerlessJobSensorTrigger)
+ assert trigger.serialized_fields == {
+ "application_id": self.app_id,
+ "job_run_id": self.job_run_id,
+ "target_states": {"RUNNING"},
+ }
+ assert trigger.waiter_delay == 10
+ assert trigger.aws_conn_id == "aws_default"
+ assert trigger.region_name == "eu-west-1"
+ assert trigger.verify is False
+ assert trigger.botocore_config == {"read_timeout": 42}
+ assert exc.value.timeout == timedelta(seconds=300)
+
+
@mock.patch("airflow.providers.amazon.aws.sensors.emr.EmrServerlessJobSensor.poke",
autospec=True)
+ def test_sensor_defer_skipped_when_poke_succeeds(self, mock_poke):
+ self.sensor.deferrable = True
+ mock_poke.return_value = True
+ self.sensor.execute(context=None)
+ mock_poke.assert_called_once()
+
+ def test_execute_complete_success(self):
+ self.sensor.execute_complete(context={}, event={"status": "success",
"value": None})
+
+ def test_execute_complete_failure(self):
+ with pytest.raises(RuntimeError, match="Error while running job"):
+ self.sensor.execute_complete(context={}, event={"status": "error",
"message": "Job failed"})
diff --git a/providers/amazon/tests/unit/amazon/aws/triggers/test_emr.py
b/providers/amazon/tests/unit/amazon/aws/triggers/test_emr.py
index 2e8ea174c1b..4e497fc52d5 100644
--- a/providers/amazon/tests/unit/amazon/aws/triggers/test_emr.py
+++ b/providers/amazon/tests/unit/amazon/aws/triggers/test_emr.py
@@ -22,7 +22,7 @@ from unittest import mock
import pytest
-from airflow.providers.amazon.aws.hooks.emr import EmrContainerHook
+from airflow.providers.amazon.aws.hooks.emr import EmrContainerHook,
EmrServerlessHook
from airflow.providers.amazon.aws.triggers.emr import (
EmrAddStepsTrigger,
EmrContainerTrigger,
@@ -30,6 +30,7 @@ from airflow.providers.amazon.aws.triggers.emr import (
EmrServerlessCancelJobsTrigger,
EmrServerlessCreateApplicationTrigger,
EmrServerlessDeleteApplicationTrigger,
+ EmrServerlessJobSensorTrigger,
EmrServerlessStartApplicationTrigger,
EmrServerlessStartJobTrigger,
EmrServerlessStopApplicationTrigger,
@@ -316,6 +317,85 @@ class TestEmrServerlessStopApplicationTrigger:
}
+class TestEmrServerlessJobSensorTrigger:
+ @staticmethod
+ def build_trigger(**kwargs):
+ trigger_kwargs = {
+ "application_id": "test_application_id",
+ "job_run_id": "test_job_run_id",
+ "target_states": {"RUNNING"},
+ "waiter_delay": 10,
+ "aws_conn_id": "aws_default",
+ **kwargs,
+ }
+ return EmrServerlessJobSensorTrigger(**trigger_kwargs)
+
+ def test_serialization(self):
+ trigger = self.build_trigger(
+ target_states=frozenset({"SUCCESS", "RUNNING"}),
+ region_name="eu-west-1",
+ verify=False,
+ botocore_config={"read_timeout": 42},
+ )
+
+ classpath, kwargs = trigger.serialize()
+
+ assert classpath ==
"airflow.providers.amazon.aws.triggers.emr.EmrServerlessJobSensorTrigger"
+ assert kwargs == {
+ "application_id": "test_application_id",
+ "job_run_id": "test_job_run_id",
+ "target_states": {"RUNNING", "SUCCESS"},
+ "waiter_delay": 10,
+ "waiter_max_attempts": sys.maxsize,
+ "aws_conn_id": "aws_default",
+ "region_name": "eu-west-1",
+ "verify": False,
+ "botocore_config": {"read_timeout": 42},
+ }
+ recreated_trigger = EmrServerlessJobSensorTrigger(**kwargs)
+ waiter_config_overrides = trigger.waiter_config_overrides
+ recreated_waiter_config_overrides =
recreated_trigger.waiter_config_overrides
+ assert waiter_config_overrides is not None
+ assert recreated_waiter_config_overrides is not None
+ assert {
+ (acceptor["expected"], acceptor["state"])
+ for acceptor in recreated_waiter_config_overrides["acceptors"]
+ } == {(acceptor["expected"], acceptor["state"]) for acceptor in
waiter_config_overrides["acceptors"]}
+
+ assert trigger.waiter_name == "serverless_job_completed"
+ assert trigger.waiter_args == {
+ "applicationId": "test_application_id",
+ "jobRunId": "test_job_run_id",
+ }
+ assert trigger.attempts == sys.maxsize
+ assert trigger.status_queries == ["jobRun.state",
"jobRun.stateDetails"]
+ hook = trigger.hook()
+ assert hook.aws_conn_id == "aws_default"
+ assert hook._region_name == "eu-west-1"
+ assert hook._verify is False
+ assert hook._config.read_timeout == 42
+
+ def test_failure_acceptors_precede_success_acceptors(self):
+ trigger = self.build_trigger(target_states={"FAILED", "RUNNING"})
+
+ waiter_config_overrides = trigger.waiter_config_overrides
+ assert waiter_config_overrides is not None
+ acceptors = waiter_config_overrides["acceptors"]
+ failure_acceptors = [acceptor for acceptor in acceptors if
acceptor["state"] == "failure"]
+ success_acceptors = [acceptor for acceptor in acceptors if
acceptor["state"] == "success"]
+
+ assert acceptors == [*failure_acceptors, *success_acceptors]
+ assert len(failure_acceptors) ==
len(EmrServerlessHook.JOB_FAILURE_STATES)
+ assert {acceptor["expected"] for acceptor in failure_acceptors} == (
+ EmrServerlessHook.JOB_FAILURE_STATES
+ )
+ assert len(success_acceptors) == 2
+ assert {acceptor["expected"] for acceptor in success_acceptors} ==
{"FAILED", "RUNNING"}
+ assert all(
+ acceptor["matcher"] == "path" and acceptor["argument"] ==
"jobRun.state" for acceptor in acceptors
+ )
+
+
class TestEmrServerlessStartJobTrigger:
def test_serialization(self):
application_id = "test_application_id"
diff --git
a/providers/amazon/tests/unit/amazon/aws/triggers/test_serialization.py
b/providers/amazon/tests/unit/amazon/aws/triggers/test_serialization.py
index 8b5149ed343..a894318abae 100644
--- a/providers/amazon/tests/unit/amazon/aws/triggers/test_serialization.py
+++ b/providers/amazon/tests/unit/amazon/aws/triggers/test_serialization.py
@@ -38,6 +38,7 @@ from airflow.providers.amazon.aws.triggers.emr import (
EmrServerlessCancelJobsTrigger,
EmrServerlessCreateApplicationTrigger,
EmrServerlessDeleteApplicationTrigger,
+ EmrServerlessJobSensorTrigger,
EmrServerlessStartApplicationTrigger,
EmrServerlessStartJobTrigger,
EmrServerlessStopApplicationTrigger,
@@ -250,6 +251,14 @@ class TestTriggersSerialization:
waiter_delay=WAITER_DELAY,
waiter_max_attempts=MAX_ATTEMPTS,
),
+ EmrServerlessJobSensorTrigger(
+ application_id=TEST_APPLICATION_ID,
+ job_run_id=TEST_JOB_ID,
+ target_states={"RUNNING", "SUCCESS"},
+ waiter_delay=WAITER_DELAY,
+ aws_conn_id=AWS_CONN_ID,
+ region_name=AWS_REGION,
+ ),
EmrServerlessStartJobTrigger(
application_id=TEST_APPLICATION_ID,
job_id=TEST_JOB_ID,