Source code for pyathena.spark.cursor

# Copyright 2024 The PyAthena authors
#
# Licensed under the MIT License.
# See LICENSE or https://opensource.org/licenses/MIT.
#
# SPDX-License-Identifier: MIT

"""Cursor that runs PySpark code in an Athena for Apache Spark session."""

from __future__ import annotations

import logging
from typing import Any, cast

from pyathena import OperationalError, ProgrammingError
from pyathena.model import AthenaCalculationExecution, AthenaCalculationExecutionStatus
from pyathena.spark.common import SparkBaseCursor, WithCalculationExecution
from pyathena.util import override

_logger = logging.getLogger(__name__)


[docs] class SparkCursor(SparkBaseCursor, WithCalculationExecution): """Cursor for executing PySpark code on Amazon Athena for Apache Spark. This cursor allows you to execute PySpark code directly on Athena's managed Spark environment. It's designed for big data processing, ETL operations, and machine learning workloads that require Spark's distributed computing capabilities. The cursor manages Spark sessions automatically and provides an interface similar to other PyAthena cursors but optimized for Spark calculations rather than SQL queries. Attributes: session_id: The Athena Spark session ID. description: The description of the current calculation. calculation_id: ID of the current calculation being executed. Example: >>> from pyathena.spark.cursor import SparkCursor >>> cursor = connection.cursor(SparkCursor) >>> >>> # Execute PySpark code >>> spark_code = ''' ... df = spark.read.table("my_database.my_table") ... result = df.groupBy("category").count() ... result.show() ... ''' >>> cursor.execute(spark_code) >>> output = cursor.get_std_out() # Configure Spark session >>> cursor = connection.cursor( ... SparkCursor, ... engine_configuration={ ... 'CoordinatorDpuSize': 1, ... 'MaxConcurrentDpus': 20, ... 'DefaultExecutorDpuSize': 1 ... } ... ) Note: Requires an Athena workgroup configured for Spark calculations. Spark sessions have associated costs and idle timeout settings. """ @property @override def calculation_execution(self) -> AthenaCalculationExecution | None: return self._calculation_execution
[docs] def get_std_out(self) -> str | None: """Get the standard output from the Spark calculation execution. Retrieves and returns the contents of the standard output generated during the Spark calculation execution, if available. Returns: The standard output as a string, or None if no output is available or the calculation has not been executed. """ if not self._calculation_execution or not self._calculation_execution.std_out_s3_uri: return None return self._read_s3_file_as_text(self._calculation_execution.std_out_s3_uri)
[docs] def get_std_error(self) -> str | None: """Get the standard error from the Spark calculation execution. Retrieves and returns the contents of the standard error generated during the Spark calculation execution, if available. This is useful for debugging failed or problematic Spark operations. Returns: The standard error as a string, or None if no error output is available or the calculation has not been executed. """ if not self._calculation_execution or not self._calculation_execution.std_error_s3_uri: return None return self._read_s3_file_as_text(self._calculation_execution.std_error_s3_uri)
[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, ) -> SparkCursor: # A failure below must not leave the previous calculation on the cursor. self._calculation_id = None self._calculation_execution = None self._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, ) self._calculation_execution = cast( AthenaCalculationExecution, self._poll(self._calculation_id) ) if self._calculation_execution.state != AthenaCalculationExecutionStatus.STATE_COMPLETED: std_error = self.get_std_error() raise OperationalError(std_error) return self
[docs] def cancel(self) -> None: """Stop the calculation that ``execute()`` last started. Raises: ProgrammingError: If no calculation ID is set. OperationalError: If the ``StopCalculationExecution`` request fails. """ if not self.calculation_id: raise ProgrammingError("CalculationExecutionId is none or empty.") self._cancel(self.calculation_id)