Source code for pyathena.spark.async_cursor

# 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)