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,

Reply via email to