# Copyright 2017 The PyAthena authors
#
# Licensed under the MIT License.
# See LICENSE or https://opensource.org/licenses/MIT.
#
# SPDX-License-Identifier: MIT
"""Asynchronous cursor that runs PySpark code in an Athena for Apache Spark session."""
import logging
from concurrent.futures import Future, ThreadPoolExecutor
from multiprocessing import cpu_count
from typing import TYPE_CHECKING, Any, cast
from pyathena.model import AthenaCalculationExecution
from pyathena.spark.common import SparkBaseCursor
from pyathena.util import override
if TYPE_CHECKING:
from pyathena.model import AthenaQueryExecution
_logger = logging.getLogger(__name__)
[docs]
class AsyncSparkCursor(SparkBaseCursor):
"""Asynchronous cursor for executing PySpark code on Amazon Athena for Apache Spark.
This cursor provides asynchronous execution of PySpark code on Athena's managed
Spark environment. It's designed for non-blocking big data processing, ETL
operations, and machine learning workloads that require Spark's distributed
computing capabilities without blocking the main thread.
Features:
- Asynchronous PySpark code execution with concurrent futures
- Non-blocking query submission and result polling
- Managed Spark sessions with configurable resources
- Access to standard output and error streams asynchronously
- Automatic session lifecycle management
- Thread pool executor for concurrent operations
Attributes:
session_id: The Athena Spark session ID.
Example:
>>> from pyathena.spark.async_cursor import AsyncSparkCursor
>>>
>>> cursor = connection.cursor(
... AsyncSparkCursor,
... engine_configuration={
... 'CoordinatorDpuSize': 1,
... 'MaxConcurrentDpus': 20
... }
... )
>>>
>>> # Execute PySpark code asynchronously
>>> spark_code = '''
... df = spark.read.table("my_database.my_table")
... result = df.groupBy("category").count()
... result.show()
... '''
>>> calculation_id, future = cursor.execute(spark_code)
>>>
>>> # Get result when ready
>>> calc_execution = future.result()
>>> stdout_future = cursor.get_std_out(calc_execution)
>>> if stdout_future:
... output = stdout_future.result()
... print(output)
Note:
Requires an Athena workgroup configured for Spark calculations.
Spark sessions have associated costs and idle timeout settings.
The cursor manages a thread pool for asynchronous operations.
"""
[docs]
def __init__(
self,
session_id: str | None = None,
description: str | None = None,
engine_configuration: dict[str, Any] | None = None,
notebook_version: str | None = None,
session_idle_timeout_minutes: int | None = None,
max_workers: int = (cpu_count() or 1) * 5,
terminate_session_on_close: bool | None = None,
**kwargs,
):
"""Initialize the cursor and start or attach to a Spark session.
Args:
session_id: ID of an existing session to use. If omitted, a new
session is started.
description: Description of a new session.
engine_configuration: Engine configuration of a new session.
notebook_version: Notebook version of a new session.
session_idle_timeout_minutes: Idle timeout of a new session in minutes.
max_workers: Maximum number of threads for asynchronous operations.
terminate_session_on_close: Whether ``close()`` terminates the session.
If None, only a session started by this cursor is terminated;
a session supplied with ``session_id`` is left running.
**kwargs: Arguments passed to ``SparkBaseCursor``.
Raises:
ValueError: If ``max_workers`` is not greater than 0.
OperationalError: If the supplied session does not exist, or the
session cannot be started or does not become idle.
"""
# Created before the session so that an invalid max_workers cannot leave
# a newly started session behind; the executor starts no threads until used.
self._max_workers = max_workers
self._executor = ThreadPoolExecutor(max_workers=max_workers)
super().__init__(
session_id=session_id,
description=description,
engine_configuration=engine_configuration,
notebook_version=notebook_version,
session_idle_timeout_minutes=session_idle_timeout_minutes,
terminate_session_on_close=terminate_session_on_close,
**kwargs,
)
[docs]
@override
def close(self, wait: bool = False) -> None:
"""Close the cursor, then shut down the executor.
The session is terminated as described in ``SparkBaseCursor.close()``.
The executor is shut down even if terminating the session fails.
Args:
wait: Whether to wait for submitted futures to finish before returning
or raising.
Raises:
OperationalError: If terminating the session fails.
"""
try:
super().close()
finally:
self._executor.shutdown(wait=wait)
[docs]
def calculation_execution(self, query_id: str) -> "Future[AthenaCalculationExecution]":
"""Get calculation execution details asynchronously.
Args:
query_id: The calculation execution ID.
Returns:
Future object containing the ``AthenaCalculationExecution``.
"""
return self._executor.submit(self._get_calculation_execution, query_id)
[docs]
def get_std_out(
self, calculation_execution: AthenaCalculationExecution
) -> "Future[str] | None":
"""Read the standard output of a calculation from S3 asynchronously.
Args:
calculation_execution: The calculation execution whose
``std_out_s3_uri`` is read.
Returns:
Future object containing the output text with leading and trailing
whitespace removed, or None if the calculation has no ``std_out_s3_uri``.
"""
if not calculation_execution.std_out_s3_uri:
return None
return self._executor.submit(
self._read_s3_file_as_text, calculation_execution.std_out_s3_uri
)
[docs]
def get_std_error(
self, calculation_execution: AthenaCalculationExecution
) -> "Future[str] | None":
"""Read the standard error output of a calculation from S3 asynchronously.
Args:
calculation_execution: The calculation execution whose
``std_error_s3_uri`` is read.
Returns:
Future object containing the error output text with leading and trailing
whitespace removed, or None if the calculation has no ``std_error_s3_uri``.
"""
if not calculation_execution.std_error_s3_uri:
return None
return self._executor.submit(
self._read_s3_file_as_text, calculation_execution.std_error_s3_uri
)
[docs]
def poll(self, query_id: str) -> "Future[AthenaCalculationExecution]":
"""Wait for a calculation to reach a terminal state asynchronously.
Args:
query_id: The calculation execution ID.
Returns:
Future object containing the calculation execution in a terminal state.
"""
return cast(
"Future[AthenaCalculationExecution]", self._executor.submit(self._poll, query_id)
)
[docs]
@override
def execute(
self,
operation: str,
parameters: dict[str, Any] | list[str] | None = None,
session_id: str | None = None,
description: str | None = None,
client_request_token: str | None = None,
work_group: str | None = None,
**kwargs,
) -> tuple[str, "Future[AthenaQueryExecution | AthenaCalculationExecution]"]:
calculation_id = self._calculate(
session_id=session_id if session_id else self._session_id,
code_block=operation,
description=description,
client_request_token=client_request_token,
)
return calculation_id, self._executor.submit(self._poll, calculation_id)
[docs]
def cancel(self, query_id: str) -> "Future[None]":
"""Stop a calculation execution asynchronously.
Args:
query_id: The calculation execution ID.
Returns:
Future object that completes when the ``StopCalculationExecution``
request has been sent.
"""
return self._executor.submit(self._cancel, query_id)