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 85321901f10 Add EmrServerlessStartSessionOperator to Amazon provider
(#70763)
85321901f10 is described below
commit 85321901f10da419bbeeecdd69a09ad872ad3eee
Author: Vincent Gromakowski <[email protected]>
AuthorDate: Tue Sep 29 18:48:34 2026 +0200
Add EmrServerlessStartSessionOperator to Amazon provider (#70763)
EMR Serverless users need an Airflow-native way to start Spark Connect
sessions without holding worker slots while AWS provisions the session.
Interactive sessions need a newer botocore than the provider's floor,
so support is checked at runtime instead of forcing an upgrade on users
who do not use sessions. Deferred completion relies only on the trigger
event, and the triggerer polls with the task's region, SSL verification
and botocore settings, so resuming does not depend on worker-local state
or the default client configuration. Dag authors who opt out of waiting
get the session identifiers right away in either execution mode.
Co-authored-by: vgkowski <[email protected]>
---
.../amazon/docs/operators/emr/emr_serverless.rst | 23 ++++
.../src/airflow/providers/amazon/aws/hooks/emr.py | 51 ++++++++
.../airflow/providers/amazon/aws/operators/emr.py | 112 +++++++++++++++++
.../airflow/providers/amazon/aws/triggers/emr.py | 48 ++++++++
.../amazon/aws/waiters/emr-serverless.json | 37 ++++++
.../amazon/aws/example_emr_serverless_session.py | 103 ++++++++++++++++
.../unit/amazon/aws/hooks/test_emr_serverless.py | 57 +++++++++
.../aws/operators/test_emr_serverless_session.py | 137 +++++++++++++++++++++
.../tests/unit/amazon/aws/triggers/test_emr.py | 60 +++++++++
9 files changed, 628 insertions(+)
diff --git a/providers/amazon/docs/operators/emr/emr_serverless.rst
b/providers/amazon/docs/operators/emr/emr_serverless.rst
index 004c825f8e2..a5e379c8970 100644
--- a/providers/amazon/docs/operators/emr/emr_serverless.rst
+++ b/providers/amazon/docs/operators/emr/emr_serverless.rst
@@ -153,6 +153,29 @@ To monitor the state of an EMR Serverless Application you
can use
:start-after: [START howto_sensor_emr_serverless_application]
:end-before: [END howto_sensor_emr_serverless_application]
+.. _howto/operator:EmrServerlessStartSessionOperator:
+
+Start an EMR Serverless interactive session
+===========================================
+
+To start an EMR Serverless interactive session that a Spark Connect client can
attach to, use
+:class:`~airflow.providers.amazon.aws.operators.emr.EmrServerlessStartSessionOperator`.
+Set ``deferrable=True`` to release the worker slot while the session warms up.
+
+.. note::
+ Interactive sessions require Amazon EMR release ``emr-7.13.0`` or later,
and the session APIs
+ are only available in ``botocore>=1.43.0``. Deferrable mode additionally
needs
+ ``aiobotocore>=3.6.0``, the first release whose ``botocore`` pin allows
1.43.0. The Amazon
+ provider keeps a lower minimum for these libraries, so install compatible
versions to use
+ interactive sessions; the operator raises a clear error at runtime if the
installed
+ ``botocore`` is too old.
+
+.. exampleinclude::
/../../amazon/tests/system/amazon/aws/example_emr_serverless_session.py
+ :language: python
+ :dedent: 4
+ :start-after: [START howto_operator_emr_serverless_start_session]
+ :end-before: [END howto_operator_emr_serverless_start_session]
+
Reference
---------
diff --git a/providers/amazon/src/airflow/providers/amazon/aws/hooks/emr.py
b/providers/amazon/src/airflow/providers/amazon/aws/hooks/emr.py
index 87c931d134f..9f134211c4a 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/hooks/emr.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/hooks/emr.py
@@ -27,6 +27,7 @@ from botocore.exceptions import ClientError
from tenacity import retry_if_exception, stop_after_attempt, wait_fixed
from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook
+from airflow.providers.amazon.aws.utils import get_botocore_version
from airflow.providers.amazon.aws.utils.waiter_with_logging import wait
from airflow.providers.common.compat.sdk import AirflowException,
AirflowNotFoundException
@@ -263,10 +264,30 @@ class EmrServerlessHook(AwsBaseHook):
APPLICATION_FAILURE_STATES = {"STOPPED", "TERMINATED"}
APPLICATION_SUCCESS_STATES = {"CREATED", "STARTED"}
+ SESSION_INTERMEDIATE_STATES = {"SUBMITTED", "STARTING"}
+ SESSION_FAILURE_STATES = {"FAILED", "TERMINATING", "TERMINATED"}
+ SESSION_SUCCESS_STATES = {"STARTED", "IDLE"}
+
+ # botocore version that first shipped the EMR Serverless interactive
session APIs.
+ # The provider keeps a lower botocore floor, so the session methods gate
on this at
+ # runtime instead of forcing every user onto a newer botocore.
+ SESSION_MIN_BOTOCORE_VERSION = (1, 43, 0)
+
def __init__(self, *args: Any, **kwargs: Any) -> None:
kwargs["client_type"] = "emr-serverless"
super().__init__(*args, **kwargs)
+ def _check_interactive_session_support(self) -> None:
+ """Raise a clear error if the installed botocore is too old for
interactive sessions."""
+ if get_botocore_version() < self.SESSION_MIN_BOTOCORE_VERSION:
+ required = ".".join(map(str, self.SESSION_MIN_BOTOCORE_VERSION))
+ installed = ".".join(map(str, get_botocore_version()))
+ raise RuntimeError(
+ f"EMR Serverless interactive sessions require botocore >=
{required}, "
+ f"but botocore {installed} is installed. Upgrade botocore (and
aiobotocore >= 3.6.0 "
+ "for deferrable mode) to use this feature."
+ )
+
def cancel_running_jobs(
self, application_id: str, waiter_config: dict | None = None,
wait_for_completion: bool = True
) -> int:
@@ -311,6 +332,36 @@ class EmrServerlessHook(AwsBaseHook):
return count
+ def start_session(
+ self,
+ application_id: str,
+ execution_role_arn: str,
+ name: str | None = None,
+ idle_timeout_minutes: int | None = None,
+ configuration_overrides: dict | None = None,
+ ) -> str:
+ """
+ Start an EMR Serverless interactive session and return its id.
+
+ :param application_id: The id of the EMR Serverless application to run
the session on.
+ :param execution_role_arn: The IAM role ARN the session assumes to
access data.
+ :param name: An optional name for the session.
+ :param idle_timeout_minutes: Auto-stop the session after this many
idle minutes.
+ :param configuration_overrides: Optional Spark/monitoring
configuration overrides.
+ """
+ self._check_interactive_session_support()
+ params: dict[str, Any] = {
+ "applicationId": application_id,
+ "executionRoleArn": execution_role_arn,
+ }
+ if name is not None:
+ params["name"] = name
+ if idle_timeout_minutes is not None:
+ params["idleTimeoutMinutes"] = idle_timeout_minutes
+ if configuration_overrides is not None:
+ params["configurationOverrides"] = configuration_overrides
+ return self.conn.start_session(**params)["sessionId"]
+
def is_connection_being_updated_exception(exception: BaseException) -> bool:
return (
diff --git a/providers/amazon/src/airflow/providers/amazon/aws/operators/emr.py
b/providers/amazon/src/airflow/providers/amazon/aws/operators/emr.py
index 4537bf23fd1..017c00d9e65 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/operators/emr.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/operators/emr.py
@@ -45,6 +45,7 @@ from airflow.providers.amazon.aws.triggers.emr import (
EmrServerlessCancelJobsTrigger,
EmrServerlessCreateApplicationTrigger,
EmrServerlessDeleteApplicationTrigger,
+ EmrServerlessSessionTrigger,
EmrServerlessStartApplicationTrigger,
EmrServerlessStartJobTrigger,
EmrServerlessStopApplicationTrigger,
@@ -1910,3 +1911,114 @@ class
EmrServerlessDeleteApplicationOperator(EmrServerlessStopApplicationOperato
if validated_event["status"] != "success":
raise AirflowException(f"Error deleting EMR Serverless
application: {validated_event}")
self.log.info("EMR serverless application %s deleted successfully",
self.application_id)
+
+
+class EmrServerlessStartSessionOperator(AwsBaseOperator[EmrServerlessHook]):
+ """
+ Start an EMR Serverless interactive session and wait until it is ready.
+
+ .. seealso::
+ For more information on how to use this operator, take a look at the
guide:
+ :ref:`howto/operator:EmrServerlessStartSessionOperator`
+
+ :param application_id: ID of the EMR Serverless application to run the
session on.
+ :param execution_role_arn: ARN of the IAM role the session assumes to
access data.
+ :param name: An optional name for the session.
+ :param idle_timeout_minutes: Auto-stop the session after this many idle
minutes.
+ :param configuration_overrides: Optional Spark/monitoring configuration
overrides.
+ :param wait_for_completion: If True, wait for the session to be ready
before returning.
+ :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
+ empty, then default boto3 configuration would be used (and must be
+ maintained on each worker node).
+ :param region_name: AWS region_name. If not specified then the default
boto3 behaviour is used.
+ :param verify: Whether or not to verify SSL certificates. See:
+
https://boto3.amazonaws.com/v1/documentation/api/latest/reference/core/session.html
+ :param waiter_max_attempts: Number of times the waiter should poll the
session to check the state.
+ :param waiter_delay: Number of seconds between polling the state of the
session.
+ :param deferrable: If True and ``wait_for_completion`` is enabled, the
operator will wait
+ asynchronously for the session to be ready. This mode requires
aiobotocore to be installed.
+ (default: False, but can be overridden in config file by setting
default_deferrable to True)
+ """
+
+ aws_hook_class = EmrServerlessHook
+ template_fields: Sequence[str] = aws_template_fields(
+ "application_id",
+ "execution_role_arn",
+ "name",
+ "idle_timeout_minutes",
+ "configuration_overrides",
+ )
+
+ def __init__(
+ self,
+ *,
+ application_id: str,
+ execution_role_arn: str,
+ name: str | None = None,
+ idle_timeout_minutes: int | None = None,
+ configuration_overrides: dict | None = None,
+ wait_for_completion: bool = True,
+ waiter_delay: int = 10,
+ waiter_max_attempts: int = 60,
+ deferrable: bool = conf.getboolean("operators", "default_deferrable",
fallback=False),
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self.application_id = application_id
+ self.execution_role_arn = execution_role_arn
+ self.name = name
+ self.idle_timeout_minutes = idle_timeout_minutes
+ self.configuration_overrides = configuration_overrides
+ self.wait_for_completion = wait_for_completion
+ self.waiter_delay = waiter_delay
+ self.waiter_max_attempts = waiter_max_attempts
+ self.deferrable = deferrable
+
+ def execute(self, context: Context) -> dict:
+ session_id = self.hook.start_session(
+ application_id=self.application_id,
+ execution_role_arn=self.execution_role_arn,
+ name=self.name,
+ idle_timeout_minutes=self.idle_timeout_minutes,
+ configuration_overrides=self.configuration_overrides,
+ )
+ self.log.info("Started EMR Serverless session %s", session_id)
+
+ if self.wait_for_completion:
+ if self.deferrable:
+ self.defer(
+ trigger=EmrServerlessSessionTrigger(
+ application_id=self.application_id,
+ session_id=session_id,
+ waiter_delay=self.waiter_delay,
+ waiter_max_attempts=self.waiter_max_attempts,
+ aws_conn_id=self.aws_conn_id,
+ region_name=self.region_name,
+ verify=self.verify,
+ botocore_config=self.botocore_config,
+ ),
+ timeout=timedelta(seconds=self.waiter_max_attempts *
self.waiter_delay),
+ method_name="execute_complete",
+ )
+ else:
+ wait(
+ waiter=self.hook.get_waiter("serverless_session_ready"),
+ waiter_delay=self.waiter_delay,
+ waiter_max_attempts=self.waiter_max_attempts,
+ args={"applicationId": self.application_id, "sessionId":
session_id},
+ failure_message="EMR Serverless session failed to start",
+ status_message="EMR Serverless session status is",
+ status_args=["session.state", "session.stateDetails"],
+ )
+ return {"application_id": self.application_id, "session_id":
session_id}
+
+ def execute_complete(self, context: Context, event: dict[str, Any] | None
= None) -> dict:
+ validated_event = validate_execute_complete_event(event)
+
+ if validated_event["status"] != "success":
+ raise RuntimeError(f"Error starting EMR Serverless session:
{validated_event}")
+ session_details = validated_event["session_details"]
+ self.log.info("EMR Serverless session %s started",
session_details["session_id"])
+ return session_details
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 8c557f494b1..24fecaf8bce 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/triggers/emr.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/triggers/emr.py
@@ -813,3 +813,51 @@ class EmrServerlessCancelJobsTrigger(AwsBaseWaiterTrigger):
def hook_instance(self) -> AwsGenericHook:
"""This property is added for backward compatibility."""
return self.hook()
+
+
+class EmrServerlessSessionTrigger(AwsBaseWaiterTrigger):
+ """
+ Poll an EMR Serverless interactive session until it reaches a ready state.
+
+ :param application_id: The ID of the EMR Serverless application.
+ :param session_id: The ID of the interactive session being polled.
+ :param waiter_delay: polling period in seconds to check for the status
+ :param waiter_max_attempts: The maximum number of attempts to be made
+ :param aws_conn_id: Reference to AWS connection id
+ :param region_name: The AWS region where the resources to watch are.
+ :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
+ """
+
+ aws_hook_class = EmrServerlessHook
+
+ def __init__(
+ self,
+ *,
+ application_id: str,
+ session_id: str,
+ waiter_delay: int = 10,
+ waiter_max_attempts: int = 60,
+ aws_conn_id: str | None = "aws_default",
+ region_name: str | None = None,
+ verify: bool | str | None = None,
+ botocore_config: dict | None = None,
+ ) -> None:
+ super().__init__(
+ serialized_fields={"application_id": application_id, "session_id":
session_id},
+ waiter_name="serverless_session_ready",
+ waiter_args={"applicationId": application_id, "sessionId":
session_id},
+ failure_message="EMR Serverless session failed to start",
+ status_message="EMR Serverless session status is",
+ status_queries=["session.state"],
+ return_key="session_details",
+ return_value={"application_id": application_id, "session_id":
session_id},
+ 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,
+ )
diff --git
a/providers/amazon/src/airflow/providers/amazon/aws/waiters/emr-serverless.json
b/providers/amazon/src/airflow/providers/amazon/aws/waiters/emr-serverless.json
index 4066109382a..ceaa231fc68 100644
---
a/providers/amazon/src/airflow/providers/amazon/aws/waiters/emr-serverless.json
+++
b/providers/amazon/src/airflow/providers/amazon/aws/waiters/emr-serverless.json
@@ -152,6 +152,43 @@
"state": "success"
}
]
+ },
+ "serverless_session_ready": {
+ "operation": "GetSession",
+ "delay": 10,
+ "maxAttempts": 60,
+ "acceptors": [
+ {
+ "matcher": "path",
+ "argument": "session.state",
+ "expected": "STARTED",
+ "state": "success"
+ },
+ {
+ "matcher": "path",
+ "argument": "session.state",
+ "expected": "IDLE",
+ "state": "success"
+ },
+ {
+ "matcher": "path",
+ "argument": "session.state",
+ "expected": "FAILED",
+ "state": "failure"
+ },
+ {
+ "matcher": "path",
+ "argument": "session.state",
+ "expected": "TERMINATING",
+ "state": "failure"
+ },
+ {
+ "matcher": "path",
+ "argument": "session.state",
+ "expected": "TERMINATED",
+ "state": "failure"
+ }
+ ]
}
}
}
diff --git
a/providers/amazon/tests/system/amazon/aws/example_emr_serverless_session.py
b/providers/amazon/tests/system/amazon/aws/example_emr_serverless_session.py
new file mode 100644
index 00000000000..ff55ad40a63
--- /dev/null
+++ b/providers/amazon/tests/system/amazon/aws/example_emr_serverless_session.py
@@ -0,0 +1,103 @@
+# 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 datetime import datetime
+
+from airflow.providers.amazon.aws.operators.emr import (
+ EmrServerlessCreateApplicationOperator,
+ EmrServerlessDeleteApplicationOperator,
+ EmrServerlessStartSessionOperator,
+ EmrServerlessStopApplicationOperator,
+)
+from airflow.providers.common.compat.sdk import DAG, TriggerRule, chain
+
+from system.amazon.aws.utils import SystemTestContextBuilder
+
+DAG_ID = "example_emr_serverless_session"
+
+# Externally fetched variables:
+ROLE_ARN_KEY = "ROLE_ARN"
+
+sys_test_context_task =
SystemTestContextBuilder().add_variable(ROLE_ARN_KEY).build()
+
+with DAG(
+ dag_id=DAG_ID,
+ schedule="@once",
+ start_date=datetime(2021, 1, 1),
+ catchup=False,
+ tags=["example", "emr-serverless"],
+) as dag:
+ test_context = sys_test_context_task()
+ role_arn = test_context[ROLE_ARN_KEY]
+
+ create_app = EmrServerlessCreateApplicationOperator(
+ task_id="create_app",
+ release_label="emr-7.13.0",
+ job_type="SPARK",
+ config={
+ "name": "session-systest",
+ # Interactive sessions (required for the StartSession API) are
only available
+ # on emr-7.13.0+ and must be explicitly enabled on the application.
+ "interactiveConfiguration": {"sessionEnabled": True},
+ },
+ )
+ application_id = create_app.output
+
+ # [START howto_operator_emr_serverless_start_session]
+ start_session = EmrServerlessStartSessionOperator(
+ task_id="start_session",
+ application_id=application_id,
+ execution_role_arn=role_arn,
+ idle_timeout_minutes=5,
+ )
+ # [END howto_operator_emr_serverless_start_session]
+
+ stop_app = EmrServerlessStopApplicationOperator(
+ task_id="stop_app",
+ application_id=application_id,
+ force_stop=True,
+ trigger_rule=TriggerRule.ALL_DONE,
+ )
+
+ delete_app = EmrServerlessDeleteApplicationOperator(
+ task_id="delete_app",
+ application_id=application_id,
+ trigger_rule=TriggerRule.ALL_DONE,
+ )
+
+ chain(
+ # TEST SETUP
+ test_context,
+ create_app,
+ # TEST BODY
+ start_session,
+ # TEST TEARDOWN
+ stop_app,
+ delete_app,
+ )
+
+ from tests_common.test_utils.watcher import watcher
+
+ # This test needs watcher in order to properly mark success/failure
+ # when "tearDown" task with trigger rule is part of the DAG
+ list(dag.tasks) >> watcher()
+
+from tests_common.test_utils.system_tests import get_test_run # noqa: E402
+
+# Needed to run the example DAG with pytest (see:
contributing-docs/testing/system_tests.rst)
+test_run = get_test_run(dag)
diff --git
a/providers/amazon/tests/unit/amazon/aws/hooks/test_emr_serverless.py
b/providers/amazon/tests/unit/amazon/aws/hooks/test_emr_serverless.py
index 75e4c5f8974..d6bf0d0794d 100644
--- a/providers/amazon/tests/unit/amazon/aws/hooks/test_emr_serverless.py
+++ b/providers/amazon/tests/unit/amazon/aws/hooks/test_emr_serverless.py
@@ -18,6 +18,8 @@ from __future__ import annotations
from unittest.mock import MagicMock, PropertyMock, patch
+import pytest
+
from airflow.providers.amazon.aws.hooks.emr import EmrServerlessHook
task_id = "test_emr_serverless_create_application_operator"
@@ -75,3 +77,58 @@ class TestEmrServerlessHook:
# nothing very interesting should happen
conn_mock.assert_called_once()
+
+
+class TestEmrServerlessHookSession:
+ @pytest.fixture(autouse=True)
+ def _supported_botocore(self):
+ # The session methods gate on botocore >= 1.43.0. Pin a supported
version so these
+ # tests exercise the boto calls regardless of the botocore installed
in CI.
+ with patch(
+ "airflow.providers.amazon.aws.hooks.emr.get_botocore_version",
+ return_value=(1, 43, 0),
+ ):
+ yield
+
+ @patch.object(EmrServerlessHook, "conn", new_callable=PropertyMock)
+ def test_start_session_minimal(self, conn_mock: MagicMock):
+ conn_mock().start_session.return_value = {"sessionId": "sess-1"}
+ hook = EmrServerlessHook(aws_conn_id="aws_default")
+
+ session_id = hook.start_session(application_id="app",
execution_role_arn="role")
+
+ assert session_id == "sess-1"
+ conn_mock().start_session.assert_called_once_with(applicationId="app",
executionRoleArn="role")
+
+ @patch.object(EmrServerlessHook, "conn", new_callable=PropertyMock)
+ def test_start_session_with_optional_params(self, conn_mock: MagicMock):
+ conn_mock().start_session.return_value = {"sessionId": "sess-2"}
+ hook = EmrServerlessHook(aws_conn_id="aws_default")
+
+ session_id = hook.start_session(
+ application_id="app",
+ execution_role_arn="role",
+ name="my-session",
+ idle_timeout_minutes=15,
+ configuration_overrides={"applicationConfiguration": []},
+ )
+
+ assert session_id == "sess-2"
+ conn_mock().start_session.assert_called_once_with(
+ applicationId="app",
+ executionRoleArn="role",
+ name="my-session",
+ idleTimeoutMinutes=15,
+ configurationOverrides={"applicationConfiguration": []},
+ )
+
+ @patch.object(EmrServerlessHook, "conn", new_callable=PropertyMock)
+ def test_start_session_gates_on_old_botocore(self, conn_mock: MagicMock):
+ hook = EmrServerlessHook(aws_conn_id="aws_default")
+ with patch(
+ "airflow.providers.amazon.aws.hooks.emr.get_botocore_version",
+ return_value=(1, 41, 0),
+ ):
+ with pytest.raises(RuntimeError, match="botocore >= 1.43.0"):
+ hook.start_session(application_id="app",
execution_role_arn="role")
+ conn_mock().start_session.assert_not_called()
diff --git
a/providers/amazon/tests/unit/amazon/aws/operators/test_emr_serverless_session.py
b/providers/amazon/tests/unit/amazon/aws/operators/test_emr_serverless_session.py
new file mode 100644
index 00000000000..0a81beb2840
--- /dev/null
+++
b/providers/amazon/tests/unit/amazon/aws/operators/test_emr_serverless_session.py
@@ -0,0 +1,137 @@
+# 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 unittest import mock
+
+import pytest
+
+from airflow.exceptions import TaskDeferred
+from airflow.providers.amazon.aws.hooks.emr import EmrServerlessHook
+from airflow.providers.amazon.aws.operators.emr import
EmrServerlessStartSessionOperator
+from airflow.providers.amazon.aws.triggers.emr import
EmrServerlessSessionTrigger
+
+APP_ID = "app-123"
+SESSION_ID = "sess-abc"
+ROLE = "arn:aws:iam::111122223333:role/emr-exec"
+REGION_NAME = "eu-west-1"
+VERIFY = "/path/to/ca.pem"
+BOTOCORE_CONFIG = {"retries": {"max_attempts": 7}}
+WAIT = "airflow.providers.amazon.aws.operators.emr.wait"
+
+
+class TestEmrServerlessStartSessionOperator:
+ @mock.patch(WAIT)
+ @mock.patch.object(EmrServerlessHook, "get_waiter")
+ @mock.patch.object(EmrServerlessHook, "start_session")
+ def test_start_and_wait(self, start_session, get_waiter, wait_mock):
+ start_session.return_value = SESSION_ID
+ op = EmrServerlessStartSessionOperator(
+ task_id="start",
+ application_id=APP_ID,
+ execution_role_arn=ROLE,
+ idle_timeout_minutes=15,
+ )
+ result = op.execute({})
+
+ start_session.assert_called_once_with(
+ application_id=APP_ID,
+ execution_role_arn=ROLE,
+ name=None,
+ idle_timeout_minutes=15,
+ configuration_overrides=None,
+ )
+ wait_mock.assert_called_once()
+ get_waiter.assert_called_once_with("serverless_session_ready")
+ assert result == {"application_id": APP_ID, "session_id": SESSION_ID}
+
+ @mock.patch(WAIT)
+ @mock.patch.object(EmrServerlessHook, "get_waiter")
+ @mock.patch.object(EmrServerlessHook, "start_session")
+ def test_no_wait(self, start_session, get_waiter, wait_mock):
+ start_session.return_value = SESSION_ID
+ op = EmrServerlessStartSessionOperator(
+ task_id="start",
+ application_id=APP_ID,
+ execution_role_arn=ROLE,
+ wait_for_completion=False,
+ )
+ op.execute({})
+ wait_mock.assert_not_called()
+
+ @mock.patch.object(EmrServerlessHook, "get_waiter")
+ @mock.patch.object(EmrServerlessHook, "start_session")
+ def test_deferrable_defers(self, start_session, get_waiter):
+ start_session.return_value = SESSION_ID
+ op = EmrServerlessStartSessionOperator(
+ task_id="start",
+ application_id=APP_ID,
+ execution_role_arn=ROLE,
+ deferrable=True,
+ region_name=REGION_NAME,
+ verify=VERIFY,
+ botocore_config=BOTOCORE_CONFIG,
+ )
+ with pytest.raises(TaskDeferred) as deferred:
+ op.execute({})
+
+ trigger = deferred.value.trigger
+ assert isinstance(trigger, EmrServerlessSessionTrigger)
+ assert trigger.return_key == "session_details"
+ assert trigger.return_value == {"application_id": APP_ID,
"session_id": SESSION_ID}
+ assert trigger.region_name == REGION_NAME
+ assert trigger.verify == VERIFY
+ assert trigger.botocore_config == BOTOCORE_CONFIG
+ get_waiter.assert_not_called()
+
+ @mock.patch.object(EmrServerlessHook, "get_waiter")
+ @mock.patch.object(EmrServerlessHook, "start_session")
+ def test_wait_for_completion_false_does_not_defer(self, start_session,
get_waiter):
+ start_session.return_value = SESSION_ID
+ op = EmrServerlessStartSessionOperator(
+ task_id="start",
+ application_id=APP_ID,
+ execution_role_arn=ROLE,
+ wait_for_completion=False,
+ deferrable=True,
+ )
+
+ result = op.execute({})
+
+ assert op.wait_for_completion is False
+ assert result == {"application_id": APP_ID, "session_id": SESSION_ID}
+ get_waiter.assert_not_called()
+
+ def test_execute_complete_success_uses_only_event_values(self):
+ op = EmrServerlessStartSessionOperator(
+ task_id="start", application_id="different-app",
execution_role_arn=ROLE
+ )
+ session_details = {"application_id": APP_ID, "session_id": SESSION_ID}
+
+ result = op.execute_complete(
+ {},
+ {"status": "success", "session_details": session_details},
+ )
+
+ assert result == session_details
+
+ def test_execute_complete_failure_raises(self):
+ op = EmrServerlessStartSessionOperator(
+ task_id="start", application_id=APP_ID, execution_role_arn=ROLE
+ )
+ with pytest.raises(RuntimeError):
+ op.execute_complete({}, {"status": "failure", "session_id":
SESSION_ID})
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 4e497fc52d5..0de41a78fb7 100644
--- a/providers/amazon/tests/unit/amazon/aws/triggers/test_emr.py
+++ b/providers/amazon/tests/unit/amazon/aws/triggers/test_emr.py
@@ -31,6 +31,7 @@ from airflow.providers.amazon.aws.triggers.emr import (
EmrServerlessCreateApplicationTrigger,
EmrServerlessDeleteApplicationTrigger,
EmrServerlessJobSensorTrigger,
+ EmrServerlessSessionTrigger,
EmrServerlessStartApplicationTrigger,
EmrServerlessStartJobTrigger,
EmrServerlessStopApplicationTrigger,
@@ -652,3 +653,62 @@ class TestEmrServerlessCancelJobsTrigger:
"waiter_max_attempts": 60,
"aws_conn_id": "aws_default",
}
+
+
+class TestEmrServerlessSessionTrigger:
+ def test_serialization(self):
+ trigger = EmrServerlessSessionTrigger(
+ application_id="test_application_id",
+ session_id="test_session_id",
+ waiter_delay=10,
+ waiter_max_attempts=60,
+ aws_conn_id="aws_default",
+ )
+ classpath, kwargs = trigger.serialize()
+ assert classpath ==
"airflow.providers.amazon.aws.triggers.emr.EmrServerlessSessionTrigger"
+ assert kwargs == {
+ "application_id": "test_application_id",
+ "session_id": "test_session_id",
+ "waiter_delay": 10,
+ "waiter_max_attempts": 60,
+ "aws_conn_id": "aws_default",
+ }
+
+ def test_serialization_with_hook_configuration(self):
+ trigger = EmrServerlessSessionTrigger(
+ application_id="test_application_id",
+ session_id="test_session_id",
+ region_name="eu-west-1",
+ verify="/path/to/ca.pem",
+ botocore_config={"retries": {"max_attempts": 7}},
+ )
+
+ _, kwargs = trigger.serialize()
+
+ assert kwargs["region_name"] == "eu-west-1"
+ assert kwargs["verify"] == "/path/to/ca.pem"
+ assert kwargs["botocore_config"] == {"retries": {"max_attempts": 7}}
+
+ def test_hook_class(self):
+ assert EmrServerlessSessionTrigger.aws_hook_class is EmrServerlessHook
+
+ def test_hook_receives_configuration(self):
+ trigger = EmrServerlessSessionTrigger(
+ application_id="test_application_id",
+ session_id="test_session_id",
+ aws_conn_id="test_conn",
+ region_name="eu-west-1",
+ verify="/path/to/ca.pem",
+ botocore_config={"retries": {"max_attempts": 7}},
+ )
+
+ with mock.patch.object(EmrServerlessSessionTrigger, "aws_hook_class")
as hook_class:
+ hook = trigger.hook()
+
+ assert hook is hook_class.return_value
+ hook_class.assert_called_once_with(
+ aws_conn_id="test_conn",
+ region_name="eu-west-1",
+ verify="/path/to/ca.pem",
+ config={"retries": {"max_attempts": 7}},
+ )