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)