This is an automated email from the ASF dual-hosted git repository.

shahar1 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 50523cd7b0c Add AthenaSparkOperator (#72081)
50523cd7b0c is described below

commit 50523cd7b0c491726c7235d177d74b37f63cc63d
Author: SameerMesiah97 <[email protected]>
AuthorDate: Tue Sep 22 05:52:14 2026 +0100

    Add AthenaSparkOperator (#72081)
---
 .../amazon/docs/operators/athena/athena_spark.rst  |  59 +++++
 providers/amazon/provider.yaml                     |   2 +
 .../airflow/providers/amazon/aws/hooks/athena.py   | 174 ++++++++++++
 .../providers/amazon/aws/operators/athena_spark.py | 204 ++++++++++++++
 .../providers/amazon/aws/waiters/athena.json       |  56 ++++
 .../airflow/providers/amazon/get_provider_info.py  |   6 +-
 .../system/amazon/aws/example_athena_spark.py      | 121 +++++++++
 .../tests/unit/amazon/aws/hooks/test_athena.py     | 225 ++++++++++++++++
 .../unit/amazon/aws/operators/test_athena_spark.py | 293 +++++++++++++++++++++
 9 files changed, 1139 insertions(+), 1 deletion(-)

diff --git a/providers/amazon/docs/operators/athena/athena_spark.rst 
b/providers/amazon/docs/operators/athena/athena_spark.rst
new file mode 100644
index 00000000000..83fe2427052
--- /dev/null
+++ b/providers/amazon/docs/operators/athena/athena_spark.rst
@@ -0,0 +1,59 @@
+.. 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.
+
+Athena Spark Operators
+======================
+
+`Amazon Athena <https://aws.amazon.com/athena/>`__ supports Apache Spark 
calculations through session-based APIs.
+This page documents the provider support for submitting and monitoring those
+calculations from Airflow.
+
+Prerequisite Tasks
+------------------
+
+.. include:: ../../_partials/prerequisite_tasks.rst
+
+Generic Parameters
+------------------
+
+.. include:: ../../_partials/generic_parameters.rst
+
+Operators
+---------
+
+.. _howto/operator:AthenaSparkOperator:
+
+Submit Spark code to an Athena session
+~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
+
+Use 
:class:`~airflow.providers.amazon.aws.operators.athena_spark.AthenaSparkOperator`
+to submit Spark code to an existing Athena Spark session.
+
+In the following example, we submit PySpark code to an existing Athena Spark
+session and wait for the calculation to complete. For more examples of how to 
use
+this operator, please see the `Sample Dag 
<https://github.com/apache/airflow/blob/|version|/providers/amazon/tests/system/amazon/aws/example_athena_spark.py>`__.
+
+.. exampleinclude:: 
/../../amazon/tests/system/amazon/aws/example_athena_spark.py
+    :language: python
+    :dedent: 4
+    :start-after: [START howto_operator_athena_spark]
+    :end-before: [END howto_operator_athena_spark]
+
+Reference
+---------
+
+* `AWS boto3 documentation for Athena calculation APIs 
<https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/athena.html>`__
diff --git a/providers/amazon/provider.yaml b/providers/amazon/provider.yaml
index 4ed477b213e..d782a3f57ad 100644
--- a/providers/amazon/provider.yaml
+++ b/providers/amazon/provider.yaml
@@ -140,6 +140,7 @@ integrations:
     how-to-guide:
       - /docs/apache-airflow-providers-amazon/operators/athena/athena_boto.rst
       - /docs/apache-airflow-providers-amazon/operators/athena/athena_sql.rst
+      - /docs/apache-airflow-providers-amazon/operators/athena/athena_spark.rst
     tags: [aws]
   - integration-name: Amazon Bedrock
     external-doc-url: https://aws.amazon.com/bedrock/
@@ -433,6 +434,7 @@ operators:
   - integration-name: Amazon Athena
     python-modules:
       - airflow.providers.amazon.aws.operators.athena
+      - airflow.providers.amazon.aws.operators.athena_spark
   - integration-name: Amazon Web Services
     python-modules:
       - airflow.providers.amazon.aws.operators.base_aws
diff --git a/providers/amazon/src/airflow/providers/amazon/aws/hooks/athena.py 
b/providers/amazon/src/airflow/providers/amazon/aws/hooks/athena.py
index 5c9621cf698..f3bfbf6b6a9 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/hooks/athena.py
+++ b/providers/amazon/src/airflow/providers/amazon/aws/hooks/athena.py
@@ -28,6 +28,7 @@ from __future__ import annotations
 from collections.abc import Collection
 from typing import TYPE_CHECKING, Any
 
+from airflow.providers.amazon.aws.exceptions import WaiterTerminalFailure
 from airflow.providers.amazon.aws.hooks.base_aws import AwsBaseHook
 from airflow.providers.amazon.aws.utils.waiter_with_logging import wait
 from airflow.providers.common.compat.sdk import AirflowException
@@ -86,6 +87,13 @@ class AthenaHook(AwsBaseHook):
         "CANCELLED",
     )
 
+    SPARK_FAILURE_STATES = (
+        "FAILED",
+        "CANCELED",
+    )
+    SPARK_SUCCESS_STATES = ("COMPLETED",)
+    SPARK_TERMINAL_STATES = SPARK_SUCCESS_STATES + SPARK_FAILURE_STATES
+
     def __init__(self, *args: Any, log_query: bool = True, **kwargs: Any) -> 
None:
         super().__init__(client_type="athena", *args, **kwargs)  # type: ignore
         self.log_query = log_query
@@ -344,3 +352,169 @@ class AthenaHook(AwsBaseHook):
         """
         self.log.info("Stopping Query with executionId - %s", 
query_execution_id)
         return 
self.get_conn().stop_query_execution(QueryExecutionId=query_execution_id)
+
+    def _get_spark_calculation_status(self, response: dict[str, Any] | None) 
-> dict[str, Any]:
+        return (response or {}).get("Status") or {}
+
+    def start_spark_calculation(
+        self,
+        *,
+        session_id: str,
+        code_block: str,
+        description: str | None = None,
+        client_request_token: str | None = None,
+    ) -> str:
+        """
+        Start an Athena Spark calculation execution.
+
+        .. seealso::
+            - 
:external+boto3:py:meth:`Athena.Client.start_calculation_execution`
+
+        :param session_id: The Athena session ID.
+        :param code_block: Spark code to execute, typically notebook-like code.
+        :param description: Optional description of the calculation. Defaults 
to None.
+        :param client_request_token: Optional idempotency token. Defaults to 
None.
+        :return: CalculationExecutionId
+        """
+        params: dict[str, Any] = {
+            "SessionId": session_id,
+            "CodeBlock": code_block,
+        }
+        if description:
+            params["Description"] = description
+
+        if client_request_token:
+            params["ClientRequestToken"] = client_request_token
+
+        if self.log_query:
+            self.log.info("Starting CalculationExecution with params:\n%s", 
query_params_to_string(params))
+
+        response = self.get_conn().start_calculation_execution(**params)
+        calculation_execution_id = response["CalculationExecutionId"]
+        self.log.info("Calculation execution id: %s", calculation_execution_id)
+        return calculation_execution_id
+
+    def get_spark_calculation_info(
+        self,
+        calculation_execution_id: str,
+        use_cache: bool = False,
+    ) -> dict[str, Any]:
+        """
+        Get information about a single Athena Spark calculation execution.
+
+        .. seealso::
+            - :external+boto3:py:meth:`Athena.Client.get_calculation_execution`
+
+        :param calculation_execution_id: CalculationExecutionId returned by 
start_spark_calculation.
+        :param use_cache: If True, use execution information cache. Defaults 
to False.
+        :return: Calculation execution response.
+        """
+        cache_key = f"calculation:{calculation_execution_id}"
+        if use_cache and cache_key in self.__query_results:
+            return self.__query_results[cache_key]
+
+        response = self.get_conn().get_calculation_execution(
+            CalculationExecutionId=calculation_execution_id,
+        )
+
+        if use_cache:
+            self.__query_results[cache_key] = response
+        return response
+
+    def check_spark_calculation_status(
+        self,
+        calculation_execution_id: str,
+        use_cache: bool = False,
+    ) -> str | None:
+        """
+        Fetch the state of a submitted Athena Spark calculation execution.
+
+        .. seealso::
+            - :external+boto3:py:meth:`Athena.Client.get_calculation_execution`
+
+        :param calculation_execution_id: CalculationExecutionId returned by 
start_spark_calculation.
+        :param use_cache: If True, use execution information cache. Defaults 
to False.
+        :return: One of valid calculation states, or *None* if the response is 
malformed.
+        """
+        response = self.get_spark_calculation_info(
+            calculation_execution_id=calculation_execution_id,
+            use_cache=use_cache,
+        )
+
+        status = self._get_spark_calculation_status(response)
+        state = status.get("State")
+
+        if state is None:
+            self.log.error("Could not parse status for calculation %s", 
calculation_execution_id)
+
+        return state
+
+    def get_spark_calculation_state_change_reason(
+        self,
+        calculation_execution_id: str,
+        use_cache: bool = False,
+    ) -> str | None:
+        """
+        Fetch the reason for an Athena Spark calculation state change, such as 
an error message.
+
+        .. seealso::
+            - :external+boto3:py:meth:`Athena.Client.get_calculation_execution`
+
+        :param calculation_execution_id: CalculationExecutionId returned by 
start_spark_calculation.
+        :param use_cache: If True, use execution information cache. Defaults 
to False.
+        :return: State change reason string, or None.
+        """
+        response = self.get_spark_calculation_info(
+            calculation_execution_id=calculation_execution_id,
+            use_cache=use_cache,
+        )
+        return 
self._get_spark_calculation_status(response).get("StateChangeReason")
+
+    def poll_spark_calculation_status(
+        self,
+        calculation_execution_id: str,
+        waiter_delay: int = 30,
+        waiter_max_attempts: int = 120,
+    ) -> str | None:
+        """
+        Poll an Athena Spark calculation until it reaches a terminal state.
+
+        :param calculation_execution_id: ID of the submitted calculation.
+        :param waiter_delay: Seconds to wait between status checks. Defaults 
to 30.
+        :param waiter_max_attempts: Maximum number of status checks. Defaults 
to 120.
+        :return: The latest calculation state, or ``None`` if the response is 
malformed.
+        """
+        try:
+            wait(
+                waiter=self.get_waiter("calculation_complete"),
+                waiter_delay=waiter_delay,
+                waiter_max_attempts=waiter_max_attempts,
+                args={"CalculationExecutionId": calculation_execution_id},
+                failure_message=(
+                    f"Error while waiting for calculation 
{calculation_execution_id} to complete"
+                ),
+                status_message=(
+                    f"Calculation execution ID {calculation_execution_id} is 
still in a non-terminal state"
+                ),
+                status_args=["Status.State"],
+            )
+
+        except WaiterTerminalFailure as error:
+            return 
self._get_spark_calculation_status(error.last_response).get("State")
+
+        return self.check_spark_calculation_status(calculation_execution_id)
+
+    def stop_spark_calculation(self, calculation_execution_id: str) -> 
dict[str, Any]:
+        """
+        Cancel the submitted Athena Spark calculation execution.
+
+        .. seealso::
+            - 
:external+boto3:py:meth:`Athena.Client.stop_calculation_execution`
+
+        :param calculation_execution_id: CalculationExecutionId returned by 
start_spark_calculation.
+        :return: Response from stop_calculation_execution.
+        """
+        self.log.info("Stopping CalculationExecution with id - %s", 
calculation_execution_id)
+        return self.get_conn().stop_calculation_execution(
+            CalculationExecutionId=calculation_execution_id,
+        )
diff --git 
a/providers/amazon/src/airflow/providers/amazon/aws/operators/athena_spark.py 
b/providers/amazon/src/airflow/providers/amazon/aws/operators/athena_spark.py
new file mode 100644
index 00000000000..b79c5749107
--- /dev/null
+++ 
b/providers/amazon/src/airflow/providers/amazon/aws/operators/athena_spark.py
@@ -0,0 +1,204 @@
+#
+# 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 Sequence
+from typing import TYPE_CHECKING, Any
+
+from airflow.providers.amazon.aws.exceptions import WaiterMaxAttemptsError
+from airflow.providers.amazon.aws.hooks.athena import AthenaHook
+from airflow.providers.amazon.aws.operators.base_aws import AwsBaseOperator
+from airflow.providers.amazon.aws.utils.mixins import aws_template_fields
+
+if TYPE_CHECKING:
+    from airflow.sdk import Context
+
+
+class AthenaSparkOperator(AwsBaseOperator[AthenaHook]):
+    """
+    Run an Apache Spark calculation in an Amazon Athena session.
+
+    Submits a calculation, such as PySpark code, via the Athena API, polls 
until
+    the calculation reaches a terminal state, and returns execution metadata.
+
+    .. seealso::
+        - :class:`airflow.providers.amazon.aws.hooks.athena.AthenaHook`
+        - `Athena for Apache Spark
+          
<https://docs.aws.amazon.com/athena/latest/ug/notebooks-spark-api-list.html>`__
+
+    :param session_id: The Athena session ID in which to run the calculation. 
(templated)
+    :param code_block: The calculation code, such as PySpark, to execute. 
(templated)
+    :param description: Optional description of the calculation. Defaults to 
None.
+    :param client_request_token: Optional idempotency token for the 
submission. Defaults to None.
+    :param waiter_delay: Seconds to wait between status checks. Defaults to 30.
+    :param waiter_max_attempts: Maximum number of polling attempts before 
timing out. Defaults to 120.
+        To limit total task time, use execution_timeout on the task as well.
+    :param log_query: Whether to log submission details. Defaults to True.
+    :param aws_conn_id: The Airflow connection used for AWS credentials. 
Defaults to ``aws_default``.
+    :param region_name: AWS region. If not set, default boto3 behavior is 
used. Defaults to None.
+    :param verify: Whether to verify SSL certificates. Defaults to None.
+    :param botocore_config: Optional botocore configuration dict. Defaults to 
None.
+    """
+
+    aws_hook_class = AthenaHook
+    ui_color = "#44b5e2"
+    template_fields: Sequence[str] = aws_template_fields("session_id", 
"code_block", "description")
+    template_ext: Sequence[str] = (".py",)
+    template_fields_renderers = {"code_block": "python"}
+
+    def __init__(
+        self,
+        *,
+        session_id: str,
+        code_block: str,
+        description: str | None = None,
+        client_request_token: str | None = None,
+        waiter_delay: int = 30,
+        waiter_max_attempts: int = 120,
+        log_query: bool = True,
+        aws_conn_id: str | None = "aws_default",
+        region_name: str | None = None,
+        verify: bool | str | None = None,
+        botocore_config: dict | None = None,
+        **kwargs: Any,
+    ) -> None:
+        super().__init__(
+            aws_conn_id=aws_conn_id,
+            region_name=region_name,
+            verify=verify,
+            botocore_config=botocore_config,
+            **kwargs,
+        )
+        self.session_id = session_id
+        self.code_block = code_block
+        self.description = description
+        self.client_request_token = client_request_token
+        self.waiter_delay = waiter_delay
+        self.waiter_max_attempts = waiter_max_attempts
+        self.log_query = log_query
+        self._calculation_execution_id: str | None = None
+
+    @property
+    def _hook_parameters(self) -> dict[str, Any]:
+        return {**super()._hook_parameters, "log_query": self.log_query}
+
+    def execute(self, context: Context) -> dict[str, Any]:
+        """Submit the Spark calculation, poll until terminal state, then 
return metadata."""
+        self.log.info("Starting Athena Spark calculation in session %s", 
self.session_id)
+
+        calculation_execution_id = self.hook.start_spark_calculation(
+            session_id=self.session_id,
+            code_block=self.code_block,
+            description=self.description,
+            client_request_token=self.client_request_token,
+        )
+        self._calculation_execution_id = calculation_execution_id
+
+        self.log.info("Calculation submitted. CalculationExecutionId: %s", 
calculation_execution_id)
+
+        try:
+            final_state = self.hook.poll_spark_calculation_status(
+                calculation_execution_id,
+                waiter_delay=self.waiter_delay,
+                waiter_max_attempts=self.waiter_max_attempts,
+            )
+        except WaiterMaxAttemptsError:
+            self._stop_calculation(calculation_execution_id)
+            raise
+
+        if final_state is None:
+            self._stop_calculation(calculation_execution_id)
+            raise RuntimeError(f"Malformed or missing status for calculation 
{calculation_execution_id}.")
+
+        if final_state not in AthenaHook.SPARK_TERMINAL_STATES:
+            self._stop_calculation(calculation_execution_id)
+            raise RuntimeError(
+                f"Polling timed out after {self.waiter_max_attempts} attempts 
for calculation "
+                f"{calculation_execution_id}."
+            )
+
+        return self._handle_terminal_state(calculation_execution_id, 
final_state)
+
+    def _stop_calculation(self, calculation_execution_id: str) -> None:
+        self.log.info("Stopping Athena Spark calculation %s", 
calculation_execution_id)
+        try:
+            self.hook.stop_spark_calculation(calculation_execution_id)
+        except Exception:
+            self.log.warning(
+                "Failed to stop Athena Spark calculation %s",
+                calculation_execution_id,
+                exc_info=True,
+            )
+
+    def _handle_terminal_state(self, calculation_execution_id: str, state: 
str) -> dict[str, Any]:
+        """Resolve terminal state: raise on failure/cancel, build and return 
metadata."""
+        reason = 
self.hook.get_spark_calculation_state_change_reason(calculation_execution_id)
+        execution_info = 
self.hook.get_spark_calculation_info(calculation_execution_id)
+        status = execution_info.get("Status", {})
+        result_info = execution_info.get("Result", {})
+
+        submission_time = status.get("SubmissionDateTime")
+        completion_time = status.get("CompletionDateTime")
+
+        result = {
+            "calculation_execution_id": calculation_execution_id,
+            "state": state,
+            "state_change_reason": reason,
+            "submission_time": str(submission_time) if submission_time else 
None,
+            "completion_time": str(completion_time) if completion_time else 
None,
+            "session_id": execution_info.get("SessionId") or self.session_id,
+            "working_directory": execution_info.get("WorkingDirectory"),
+            "stdout_s3_uri": result_info.get("StdOutS3Uri"),
+            "stderr_s3_uri": result_info.get("StdErrorS3Uri"),
+            "result_s3_uri": result_info.get("ResultS3Uri"),
+            "result_type": result_info.get("ResultType"),
+        }
+
+        if state in AthenaHook.SPARK_FAILURE_STATES:
+            self.log.error(
+                "Calculation failed. CalculationExecutionId: %s, state: %s, 
reason: %s",
+                calculation_execution_id,
+                state,
+                reason,
+            )
+            raise RuntimeError(
+                f"Athena Spark calculation ended in {state}. "
+                f"CalculationExecutionId: {calculation_execution_id}. "
+                f"Reason: {reason or 'No reason provided.'}"
+            )
+
+        if state not in AthenaHook.SPARK_SUCCESS_STATES:
+            raise RuntimeError(
+                f"Unexpected terminal state: {state} for calculation 
{calculation_execution_id}. "
+                f"Expected one of: {', 
'.join(AthenaHook.SPARK_TERMINAL_STATES)}."
+            )
+
+        self.log.info(
+            "Calculation completed successfully. CalculationExecutionId: %s",
+            calculation_execution_id,
+        )
+        return result
+
+    def on_kill(self) -> None:
+        """Request cancellation of the calculation when the task is killed."""
+        if self._calculation_execution_id:
+            self.log.info(
+                "Received kill signal for Athena Spark calculation %s", 
self._calculation_execution_id
+            )
+            self._stop_calculation(self._calculation_execution_id)
diff --git 
a/providers/amazon/src/airflow/providers/amazon/aws/waiters/athena.json 
b/providers/amazon/src/airflow/providers/amazon/aws/waiters/athena.json
index db68ce32f4b..00cf7142f8c 100644
--- a/providers/amazon/src/airflow/providers/amazon/aws/waiters/athena.json
+++ b/providers/amazon/src/airflow/providers/amazon/aws/waiters/athena.json
@@ -25,6 +25,62 @@
                     "argument": "QueryExecution.Status.State"
                 }
             ]
+        },
+        "calculation_complete": {
+        "operation": "GetCalculationExecution",
+        "delay": 30,
+        "maxAttempts": 120,
+        "acceptors": [
+                {
+                "expected": "COMPLETED",
+                "matcher": "path",
+                "state": "success",
+                "argument": "Status.State"
+                },
+                {
+                "expected": "FAILED",
+                "matcher": "path",
+                "state": "failure",
+                "argument": "Status.State"
+                },
+                {
+                "expected": "CANCELED",
+                "matcher": "path",
+                "state": "failure",
+                "argument": "Status.State"
+                }
+            ]
+        },
+        "session_idle": {
+        "operation": "GetSession",
+        "delay": 10,
+        "maxAttempts": 60,
+        "acceptors": [
+                {
+                "expected": "IDLE",
+                "matcher": "path",
+                "state": "success",
+                "argument": "Status.State"
+                },
+                {
+                "expected": "TERMINATED",
+                "matcher": "path",
+                "state": "failure",
+                "argument": "Status.State"
+                },
+                {
+                "expected": "DEGRADED",
+                "matcher": "path",
+                "state": "failure",
+                "argument": "Status.State"
+                },
+                {
+                "expected": "FAILED",
+                "matcher": "path",
+                "state": "failure",
+                "argument": "Status.State"
+                }
+            ]
         }
     }
 }
diff --git a/providers/amazon/src/airflow/providers/amazon/get_provider_info.py 
b/providers/amazon/src/airflow/providers/amazon/get_provider_info.py
index 82ed31ca7a8..80e314ef7f1 100644
--- a/providers/amazon/src/airflow/providers/amazon/get_provider_info.py
+++ b/providers/amazon/src/airflow/providers/amazon/get_provider_info.py
@@ -34,6 +34,7 @@ def get_provider_info():
                 "how-to-guide": [
                     
"/docs/apache-airflow-providers-amazon/operators/athena/athena_boto.rst",
                     
"/docs/apache-airflow-providers-amazon/operators/athena/athena_sql.rst",
+                    
"/docs/apache-airflow-providers-amazon/operators/athena/athena_spark.rst",
                 ],
                 "tags": ["aws"],
             },
@@ -396,7 +397,10 @@ def get_provider_info():
         "operators": [
             {
                 "integration-name": "Amazon Athena",
-                "python-modules": 
["airflow.providers.amazon.aws.operators.athena"],
+                "python-modules": [
+                    "airflow.providers.amazon.aws.operators.athena",
+                    "airflow.providers.amazon.aws.operators.athena_spark",
+                ],
             },
             {
                 "integration-name": "Amazon Web Services",
diff --git a/providers/amazon/tests/system/amazon/aws/example_athena_spark.py 
b/providers/amazon/tests/system/amazon/aws/example_athena_spark.py
new file mode 100644
index 00000000000..d7ccb624127
--- /dev/null
+++ b/providers/amazon/tests/system/amazon/aws/example_athena_spark.py
@@ -0,0 +1,121 @@
+# 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
+
+import boto3
+
+from airflow.providers.amazon.aws.hooks.athena import AthenaHook
+from airflow.providers.amazon.aws.operators.athena_spark import 
AthenaSparkOperator
+
+from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS
+
+if AIRFLOW_V_3_0_PLUS:
+    from airflow.sdk import DAG, chain, task
+else:
+    # Airflow 2 path
+    from airflow.decorators import task  # type: ignore[attr-defined,no-redef]
+    from airflow.models.baseoperator import chain  # type: 
ignore[attr-defined,no-redef]
+    from airflow.models.dag import DAG  # type: 
ignore[attr-defined,no-redef,assignment]
+
+try:
+    from airflow.sdk import TriggerRule
+except ImportError:
+    # Compatibility for Airflow < 3.1
+    from airflow.utils.trigger_rule import TriggerRule  # type: 
ignore[no-redef,attr-defined]
+
+from system.amazon.aws.utils import SystemTestContextBuilder
+
+DAG_ID = "example_athena_spark"
+
+# The Spark workgroup is preconfigured test infrastructure; this DAG creates 
only the session.
+# Test runners can override the default by exporting ATHENA_SPARK_WORK_GROUP.
+ATHENA_SPARK_WORK_GROUP_KEY = "ATHENA_SPARK_WORK_GROUP"
+
+sys_test_context_task = 
SystemTestContextBuilder().add_variable(ATHENA_SPARK_WORK_GROUP_KEY).build()
+
+
+@task
+def start_athena_spark_session(work_group: str) -> str:
+    client = boto3.client("athena")
+    response = client.start_session(
+        WorkGroup=work_group,
+        EngineConfiguration={"MaxConcurrentDpus": 20},
+    )
+    return response["SessionId"]
+
+
+@task
+def wait_for_athena_spark_session(session_id: str) -> str:
+    AthenaHook().get_waiter("session_idle").wait(
+        SessionId=session_id,
+        WaiterConfig={"Delay": 10, "MaxAttempts": 60},
+    )
+    return session_id
+
+
+@task(trigger_rule=TriggerRule.ALL_DONE)
+def stop_athena_spark_session(session_id: str) -> None:
+    client = boto3.client("athena")
+    client.terminate_session(SessionId=session_id)
+
+
+with DAG(
+    dag_id=DAG_ID,
+    schedule="@once",
+    start_date=datetime(2021, 1, 1),
+    catchup=False,
+) as dag:
+    test_context = sys_test_context_task()
+    athena_spark_work_group = test_context[ATHENA_SPARK_WORK_GROUP_KEY]
+
+    session_id = start_athena_spark_session(athena_spark_work_group)
+    idle_session_id = wait_for_athena_spark_session(session_id)
+
+    # [START howto_operator_athena_spark]
+    run_spark_calculation = AthenaSparkOperator(
+        task_id="run_spark_calculation",
+        session_id=idle_session_id,
+        code_block="print('hello from athena spark')",
+        waiter_delay=30,
+        waiter_max_attempts=120,
+    )
+    # [END howto_operator_athena_spark]
+
+    stop_session = stop_athena_spark_session(session_id)
+
+    chain(
+        # TEST SETUP
+        test_context,
+        session_id,
+        idle_session_id,
+        # TEST BODY
+        run_spark_calculation,
+        # TEST TEARDOWN
+        stop_session,
+    )
+
+    from tests_common.test_utils.watcher import watcher
+
+    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_athena.py 
b/providers/amazon/tests/unit/amazon/aws/hooks/test_athena.py
index e743e831873..4e09bbec92a 100644
--- a/providers/amazon/tests/unit/amazon/aws/hooks/test_athena.py
+++ b/providers/amazon/tests/unit/amazon/aws/hooks/test_athena.py
@@ -21,6 +21,10 @@ from unittest import mock
 import pytest
 from moto import mock_aws
 
+from airflow.providers.amazon.aws.exceptions import (
+    WaiterMaxAttemptsError,
+    WaiterTerminalFailure,
+)
 from airflow.providers.amazon.aws.hooks.athena import (
     MULTI_LINE_QUERY_LOG_PREFIX,
     AthenaHook,
@@ -40,6 +44,10 @@ MOCK_DATA = {
     "query_execution_id": "eac427d0-1c6d-4dfb-96aa-2835d3ac6595",
     "next_token_id": "eac427d0-1c6d-4dfb-96aa-2835d3ac6595",
     "max_items": 1000,
+    "code_block": "print('hello spark')",
+    "calculation_execution_id": "calc-123456",
+    "session_id": "session-123456",
+    "description": "spark-calc",
 }
 
 mock_query_context = {"Database": MOCK_DATA["database"]}
@@ -58,6 +66,11 @@ MOCK_QUERY_EXECUTION_OUTPUT = {
     }
 }
 
+MOCK_CALCULATION_EXECUTION = {"CalculationExecutionId": 
MOCK_DATA["calculation_execution_id"]}
+
+MOCK_RUNNING_CALC_EXECUTION = {"Status": {"State": "RUNNING"}}
+MOCK_SUCCEEDED_CALC_EXECUTION = {"Status": {"State": "COMPLETED"}}
+
 
 @mock_aws
 class TestAthenaHook:
@@ -314,3 +327,215 @@ class TestAthenaHook:
         assert result.count("\n") == len(params.keys()) + num_query_lines
         # All lines except the first line of the multiline query log message 
get the double prefix/indent.
         assert result.count(MULTI_LINE_QUERY_LOG_PREFIX) == (num_query_lines * 
2) - 1
+
+    @mock.patch.object(AthenaHook, "get_conn")
+    @pytest.mark.parametrize(
+        ("optional_kwargs", "expected_optional_params"),
+        [
+            ({}, {}),
+            (
+                {
+                    "description": MOCK_DATA["description"],
+                    "client_request_token": MOCK_DATA["client_request_token"],
+                },
+                {
+                    "Description": MOCK_DATA["description"],
+                    "ClientRequestToken": MOCK_DATA["client_request_token"],
+                },
+            ),
+        ],
+    )
+    def test_hook_start_spark_calculation(self, mock_conn, optional_kwargs, 
expected_optional_params):
+        mock_conn.return_value.start_calculation_execution.return_value = 
MOCK_CALCULATION_EXECUTION
+
+        result = self.athena.start_spark_calculation(
+            session_id=MOCK_DATA["session_id"],
+            code_block=MOCK_DATA["code_block"],
+            **optional_kwargs,
+        )
+
+        
mock_conn.return_value.start_calculation_execution.assert_called_once_with(
+            SessionId=MOCK_DATA["session_id"],
+            CodeBlock=MOCK_DATA["code_block"],
+            **expected_optional_params,
+        )
+        assert result == MOCK_DATA["calculation_execution_id"]
+
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_hook_get_spark_calculation_info(self, mock_conn):
+        mock_conn.return_value.get_calculation_execution.return_value = 
MOCK_SUCCEEDED_CALC_EXECUTION
+
+        result = self.athena.get_spark_calculation_info(
+            calculation_execution_id=MOCK_DATA["calculation_execution_id"]
+        )
+
+        
mock_conn.return_value.get_calculation_execution.assert_called_once_with(
+            CalculationExecutionId=MOCK_DATA["calculation_execution_id"]
+        )
+        assert result == MOCK_SUCCEEDED_CALC_EXECUTION
+
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_hook_get_spark_calculation_info_uses_cache(self, mock_conn):
+        mock_conn.return_value.get_calculation_execution.return_value = 
MOCK_SUCCEEDED_CALC_EXECUTION
+
+        self.athena.get_spark_calculation_info(
+            calculation_execution_id=MOCK_DATA["calculation_execution_id"],
+            use_cache=True,
+        )
+        result = self.athena.get_spark_calculation_info(
+            calculation_execution_id=MOCK_DATA["calculation_execution_id"],
+            use_cache=True,
+        )
+
+        
mock_conn.return_value.get_calculation_execution.assert_called_once_with(
+            CalculationExecutionId=MOCK_DATA["calculation_execution_id"]
+        )
+        assert result == MOCK_SUCCEEDED_CALC_EXECUTION
+
+    @mock.patch.object(AthenaHook, "get_conn")
+    @pytest.mark.parametrize(
+        ("response", "expected_state"),
+        [
+            ({"Status": {"State": "RUNNING"}}, "RUNNING"),
+            ({"Status": {"State": "COMPLETED"}}, "COMPLETED"),
+            (None, None),
+            ({}, None),
+            ({"Status": None}, None),
+            ({"Status": {}}, None),
+        ],
+    )
+    def test_hook_check_spark_calculation_status(self, mock_conn, response, 
expected_state):
+        mock_conn.return_value.get_calculation_execution.return_value = 
response
+
+        state = self.athena.check_spark_calculation_status(
+            calculation_execution_id=MOCK_DATA["calculation_execution_id"]
+        )
+
+        assert state == expected_state
+
+    @mock.patch.object(AthenaHook, "get_conn")
+    @pytest.mark.parametrize(
+        ("response", "expected_reason"),
+        [
+            (
+                {
+                    "Status": {
+                        "State": "FAILED",
+                        "StateChangeReason": "Calculation failed",
+                    }
+                },
+                "Calculation failed",
+            ),
+            ({"Status": {"State": "FAILED"}}, None),
+            ({"Status": {"State": "COMPLETED"}}, None),
+            (None, None),
+            ({}, None),
+            ({"Status": None}, None),
+            ({"Status": {}}, None),
+        ],
+    )
+    def test_hook_get_spark_calculation_state_change_reason(self, mock_conn, 
response, expected_reason):
+        mock_conn.return_value.get_calculation_execution.return_value = 
response
+
+        reason = self.athena.get_spark_calculation_state_change_reason(
+            calculation_execution_id=MOCK_DATA["calculation_execution_id"]
+        )
+
+        assert reason == expected_reason
+
+    @mock.patch("airflow.providers.amazon.aws.hooks.athena.wait")
+    @mock.patch.object(AthenaHook, "get_waiter")
+    @mock.patch.object(
+        AthenaHook,
+        "check_spark_calculation_status",
+        return_value="COMPLETED",
+    )
+    def test_hook_poll_spark_calculation_status(
+        self,
+        mock_check_status,
+        mock_get_waiter,
+        mock_wait,
+    ):
+        result = self.athena.poll_spark_calculation_status(
+            calculation_execution_id=MOCK_DATA["calculation_execution_id"],
+            waiter_delay=5,
+            waiter_max_attempts=10,
+        )
+
+        mock_get_waiter.assert_called_once_with("calculation_complete")
+        mock_wait.assert_called_once_with(
+            waiter=mock_get_waiter.return_value,
+            waiter_delay=5,
+            waiter_max_attempts=10,
+            args={
+                "CalculationExecutionId": 
MOCK_DATA["calculation_execution_id"],
+            },
+            failure_message=(
+                f"Error while waiting for calculation 
{MOCK_DATA['calculation_execution_id']} to complete"
+            ),
+            status_message=(
+                f"Calculation execution ID "
+                f"{MOCK_DATA['calculation_execution_id']} is still in a 
non-terminal state"
+            ),
+            status_args=["Status.State"],
+        )
+        
mock_check_status.assert_called_once_with(MOCK_DATA["calculation_execution_id"])
+        assert result == "COMPLETED"
+
+    @mock.patch("airflow.providers.amazon.aws.hooks.athena.wait")
+    @mock.patch.object(AthenaHook, "get_waiter")
+    @mock.patch.object(AthenaHook, "check_spark_calculation_status")
+    @pytest.mark.parametrize("state", ["FAILED", "CANCELED"])
+    def test_hook_poll_spark_calculation_status_returns_terminal_failure_state(
+        self,
+        mock_check_status,
+        mock_get_waiter,
+        mock_wait,
+        state,
+    ):
+        mock_wait.side_effect = WaiterTerminalFailure(
+            "Athena Spark calculation failed",
+            last_response={"Status": {"State": state}},
+        )
+
+        result = self.athena.poll_spark_calculation_status(
+            calculation_execution_id=MOCK_DATA["calculation_execution_id"],
+            waiter_delay=5,
+            waiter_max_attempts=10,
+        )
+
+        mock_get_waiter.assert_called_once_with("calculation_complete")
+        mock_wait.assert_called_once()
+        mock_check_status.assert_not_called()
+        assert result == state
+
+    @mock.patch(
+        "airflow.providers.amazon.aws.hooks.athena.wait",
+        side_effect=WaiterMaxAttemptsError("Waiter error: max attempts 
reached"),
+    )
+    @mock.patch.object(AthenaHook, "get_waiter")
+    @mock.patch.object(AthenaHook, "check_spark_calculation_status")
+    def test_hook_poll_spark_calculation_status_raises_after_max_attempts(
+        self,
+        mock_check_status,
+        mock_get_waiter,
+        mock_wait,
+    ):
+        with pytest.raises(WaiterMaxAttemptsError, match="max attempts 
reached"):
+            self.athena.poll_spark_calculation_status(
+                calculation_execution_id=MOCK_DATA["calculation_execution_id"],
+                waiter_delay=0,
+                waiter_max_attempts=1,
+            )
+
+        mock_get_waiter.assert_called_once_with("calculation_complete")
+        mock_wait.assert_called_once()
+        mock_check_status.assert_not_called()
+
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_hook_stop_spark_calculation(self, mock_conn):
+        
self.athena.stop_spark_calculation(calculation_execution_id=MOCK_DATA["calculation_execution_id"])
+
+        
mock_conn.return_value.stop_calculation_execution.assert_called_once_with(
+            CalculationExecutionId=MOCK_DATA["calculation_execution_id"]
+        )
diff --git 
a/providers/amazon/tests/unit/amazon/aws/operators/test_athena_spark.py 
b/providers/amazon/tests/unit/amazon/aws/operators/test_athena_spark.py
new file mode 100644
index 00000000000..97d14ba15c0
--- /dev/null
+++ b/providers/amazon/tests/unit/amazon/aws/operators/test_athena_spark.py
@@ -0,0 +1,293 @@
+#
+# 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 moto import mock_aws
+
+from airflow.models import DAG
+from airflow.providers.amazon.aws.exceptions import WaiterMaxAttemptsError
+from airflow.providers.amazon.aws.hooks.athena import AthenaHook
+from airflow.providers.amazon.aws.operators.athena_spark import 
AthenaSparkOperator
+
+from tests_common.test_utils.compat import timezone
+from unit.amazon.aws.utils.test_template_fields import validate_template_fields
+
+TEST_DAG_ID = "unit_tests"
+DEFAULT_DATE = timezone.datetime(2018, 1, 1)
+ATHENA_CALCULATION_ID = "calc-exec-123"
+
+MOCK_DATA = {
+    "task_id": "test_athena_spark_operator",
+    "session_id": "session-456",
+    "code_block": "1 + 1",
+    "description": "Test Spark calculation",
+    "client_request_token": "eac427d0-1c6d-4dfb-96aa-2835d3ac6595",
+}
+
+
+def _calculation_info(state: str, submission_time=None, completion_time=None):
+    return {
+        "CalculationExecutionId": ATHENA_CALCULATION_ID,
+        "SessionId": MOCK_DATA["session_id"],
+        "WorkingDirectory": "s3://test-bucket/spark/",
+        "Status": {
+            "State": state,
+            "SubmissionDateTime": submission_time,
+            "CompletionDateTime": completion_time,
+        },
+        "Result": {
+            "StdOutS3Uri": "s3://test-bucket/spark/stdout",
+            "StdErrorS3Uri": "s3://test-bucket/spark/stderr",
+            "ResultS3Uri": "s3://test-bucket/spark/results",
+            "ResultType": "application/vnd.aws.athena.v1+json",
+        },
+    }
+
+
+@mock_aws
+class TestAthenaSparkOperator:
+    @pytest.fixture(autouse=True)
+    def _setup_test_cases(self):
+        args = {
+            "owner": "airflow",
+            "start_date": DEFAULT_DATE,
+        }
+
+        self.dag = DAG(TEST_DAG_ID, default_args=args, schedule="@once")
+        self.default_op_kwargs = dict(
+            task_id=MOCK_DATA["task_id"],
+            session_id=MOCK_DATA["session_id"],
+            code_block=MOCK_DATA["code_block"],
+            client_request_token=MOCK_DATA["client_request_token"],
+            waiter_delay=0,
+            waiter_max_attempts=3,
+        )
+        self.athena = AthenaSparkOperator(**self.default_op_kwargs, 
aws_conn_id=None, dag=self.dag)
+
+    def test_base_aws_op_attributes(self):
+        op = AthenaSparkOperator(**self.default_op_kwargs)
+        assert op.hook.aws_conn_id == "aws_default"
+        assert op.hook._region_name is None
+        assert op.hook._verify is None
+        assert op.hook._config is None
+        assert op.hook.log_query is True
+
+        op = AthenaSparkOperator(
+            **self.default_op_kwargs,
+            aws_conn_id="aws-test-custom-conn",
+            region_name="eu-west-1",
+            verify=False,
+            botocore_config={"read_timeout": 42},
+            log_query=False,
+        )
+        assert op.hook.aws_conn_id == "aws-test-custom-conn"
+        assert op.hook._region_name == "eu-west-1"
+        assert op.hook._verify is False
+        assert op.hook._config is not None
+        assert op.hook._config.read_timeout == 42
+        assert op.hook.log_query is False
+
+    def test_init(self):
+        assert self.athena.task_id == MOCK_DATA["task_id"]
+        assert self.athena.session_id == MOCK_DATA["session_id"]
+        assert self.athena.code_block == MOCK_DATA["code_block"]
+        assert self.athena.client_request_token == 
MOCK_DATA["client_request_token"]
+        assert self.athena.waiter_delay == 0
+        assert self.athena.waiter_max_attempts == 3
+        assert self.athena._calculation_execution_id is None
+
+    @mock.patch.object(AthenaHook, "get_spark_calculation_info")
+    @mock.patch.object(AthenaHook, 
"get_spark_calculation_state_change_reason", return_value=None)
+    @mock.patch.object(AthenaHook, "poll_spark_calculation_status", 
return_value="COMPLETED")
+    @mock.patch.object(AthenaHook, "start_spark_calculation", 
return_value=ATHENA_CALCULATION_ID)
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_execute_success(
+        self,
+        mock_conn,
+        mock_start_spark_calculation,
+        mock_poll_spark_calculation_status,
+        mock_get_spark_calculation_state_change_reason,
+        mock_get_spark_calculation_info,
+    ):
+        mock_get_spark_calculation_info.return_value = 
_calculation_info("COMPLETED")
+
+        result = self.athena.execute({})
+
+        mock_start_spark_calculation.assert_called_once_with(
+            session_id=MOCK_DATA["session_id"],
+            code_block=MOCK_DATA["code_block"],
+            description=None,
+            client_request_token=MOCK_DATA["client_request_token"],
+        )
+
+        mock_poll_spark_calculation_status.assert_called_once_with(
+            ATHENA_CALCULATION_ID,
+            waiter_delay=self.athena.waiter_delay,
+            waiter_max_attempts=self.athena.waiter_max_attempts,
+        )
+        
mock_get_spark_calculation_state_change_reason.assert_called_once_with(ATHENA_CALCULATION_ID)
+        
mock_get_spark_calculation_info.assert_called_once_with(ATHENA_CALCULATION_ID)
+
+        assert result["calculation_execution_id"] == ATHENA_CALCULATION_ID
+        assert result["state"] == "COMPLETED"
+        assert result["session_id"] == MOCK_DATA["session_id"]
+        assert result["working_directory"] == "s3://test-bucket/spark/"
+        assert result["stdout_s3_uri"] == "s3://test-bucket/spark/stdout"
+        assert result["stderr_s3_uri"] == "s3://test-bucket/spark/stderr"
+        assert result["result_s3_uri"] == "s3://test-bucket/spark/results"
+        assert result["result_type"] == "application/vnd.aws.athena.v1+json"
+
+    @mock.patch.object(AthenaHook, "get_spark_calculation_info")
+    @mock.patch.object(AthenaHook, 
"get_spark_calculation_state_change_reason", return_value="Job failed")
+    @mock.patch.object(AthenaHook, "poll_spark_calculation_status", 
return_value="FAILED")
+    @mock.patch.object(AthenaHook, "start_spark_calculation", 
return_value=ATHENA_CALCULATION_ID)
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_execute_failure(
+        self,
+        mock_conn,
+        mock_start_spark_calculation,
+        mock_poll_spark_calculation_status,
+        mock_get_spark_calculation_state_change_reason,
+        mock_get_spark_calculation_info,
+    ):
+        mock_get_spark_calculation_info.return_value = 
_calculation_info("FAILED")
+
+        with pytest.raises(RuntimeError):
+            self.athena.execute({})
+
+        mock_start_spark_calculation.assert_called_once_with(
+            session_id=MOCK_DATA["session_id"],
+            code_block=MOCK_DATA["code_block"],
+            description=None,
+            client_request_token=MOCK_DATA["client_request_token"],
+        )
+
+        mock_poll_spark_calculation_status.assert_called_once_with(
+            ATHENA_CALCULATION_ID,
+            waiter_delay=self.athena.waiter_delay,
+            waiter_max_attempts=self.athena.waiter_max_attempts,
+        )
+
+        assert mock_get_spark_calculation_state_change_reason.call_count == 1
+
+    @mock.patch.object(AthenaHook, "get_spark_calculation_info")
+    @mock.patch.object(AthenaHook, 
"get_spark_calculation_state_change_reason", return_value="Canceled")
+    @mock.patch.object(AthenaHook, "poll_spark_calculation_status", 
return_value="CANCELED")
+    @mock.patch.object(AthenaHook, "start_spark_calculation", 
return_value=ATHENA_CALCULATION_ID)
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_execute_cancelled(
+        self,
+        mock_conn,
+        mock_start_spark_calculation,
+        mock_poll_spark_calculation_status,
+        mock_get_spark_calculation_state_change_reason,
+        mock_get_spark_calculation_info,
+    ):
+        mock_get_spark_calculation_info.return_value = 
_calculation_info("CANCELED")
+
+        with pytest.raises(RuntimeError):
+            self.athena.execute({})
+
+        mock_start_spark_calculation.assert_called_once_with(
+            session_id=MOCK_DATA["session_id"],
+            code_block=MOCK_DATA["code_block"],
+            description=None,
+            client_request_token=MOCK_DATA["client_request_token"],
+        )
+
+        mock_poll_spark_calculation_status.assert_called_once_with(
+            ATHENA_CALCULATION_ID,
+            waiter_delay=self.athena.waiter_delay,
+            waiter_max_attempts=self.athena.waiter_max_attempts,
+        )
+
+        assert mock_get_spark_calculation_state_change_reason.call_count == 1
+
+    @mock.patch.object(AthenaHook, "stop_spark_calculation")
+    @mock.patch.object(
+        AthenaHook,
+        "poll_spark_calculation_status",
+        side_effect=WaiterMaxAttemptsError("Waiter error: max attempts 
reached"),
+    )
+    @mock.patch.object(AthenaHook, "start_spark_calculation", 
return_value=ATHENA_CALCULATION_ID)
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_execute_timeout(
+        self,
+        mock_conn,
+        mock_start_spark_calculation,
+        mock_poll_spark_calculation_status,
+        mock_stop_spark_calculation,
+    ):
+        with pytest.raises(WaiterMaxAttemptsError, match="max attempts 
reached"):
+            self.athena.execute({})
+
+        mock_start_spark_calculation.assert_called_once_with(
+            session_id=MOCK_DATA["session_id"],
+            code_block=MOCK_DATA["code_block"],
+            description=None,
+            client_request_token=MOCK_DATA["client_request_token"],
+        )
+
+        mock_poll_spark_calculation_status.assert_called_once_with(
+            ATHENA_CALCULATION_ID,
+            waiter_delay=self.athena.waiter_delay,
+            waiter_max_attempts=self.athena.waiter_max_attempts,
+        )
+
+        
mock_stop_spark_calculation.assert_called_once_with(ATHENA_CALCULATION_ID)
+
+    @mock.patch.object(AthenaHook, "poll_spark_calculation_status", 
return_value=None)
+    @mock.patch.object(AthenaHook, "start_spark_calculation", 
return_value=ATHENA_CALCULATION_ID)
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_execute_malformed_status(
+        self,
+        mock_conn,
+        mock_start_spark_calculation,
+        mock_poll_spark_calculation_status,
+    ):
+        with pytest.raises(RuntimeError, match="Malformed or missing status"):
+            self.athena.execute({})
+
+        mock_start_spark_calculation.assert_called_once_with(
+            session_id=MOCK_DATA["session_id"],
+            code_block=MOCK_DATA["code_block"],
+            description=None,
+            client_request_token=MOCK_DATA["client_request_token"],
+        )
+
+    @mock.patch.object(AthenaHook, "stop_spark_calculation")
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_on_kill_calls_stop_spark_calculation(self, mock_conn, 
mock_stop_spark_calculation):
+        self.athena._calculation_execution_id = ATHENA_CALCULATION_ID
+
+        self.athena.on_kill()
+
+        
mock_stop_spark_calculation.assert_called_once_with(ATHENA_CALCULATION_ID)
+
+    @mock.patch.object(AthenaHook, "stop_spark_calculation")
+    @mock.patch.object(AthenaHook, "get_conn")
+    def test_on_kill_no_op_when_no_calculation_execution_id(self, mock_conn, 
mock_stop_spark_calculation):
+        self.athena.on_kill()
+
+        mock_stop_spark_calculation.assert_not_called()
+
+    def test_template_fields(self):
+        validate_template_fields(self.athena)

Reply via email to