# Copyright 2017 The PyAthena authors
#
# Licensed under the MIT License.
# See LICENSE or https://opensource.org/licenses/MIT.
#
# SPDX-License-Identifier: MIT
"""Native asyncio cursor that runs PySpark code in an Athena for Apache Spark session."""
from __future__ import annotations
import asyncio
import logging
import uuid
from typing import Any, cast
from pyathena.aio.util import async_retry_api_call
from pyathena.error import DatabaseError, NotSupportedError, OperationalError, ProgrammingError
from pyathena.model import (
AthenaCalculationExecution,
AthenaCalculationExecutionStatus,
AthenaQueryExecution,
)
from pyathena.spark.common import SparkBaseCursor, WithCalculationExecution
from pyathena.util import override, parse_output_location
_logger = logging.getLogger(__name__)
[docs]
class AioSparkCursor(SparkBaseCursor, WithCalculationExecution):
"""Native asyncio cursor for executing PySpark code on Athena.
Overrides post-init I/O methods of ``SparkBaseCursor`` with async
equivalents. Session management (``_exists_session``,
``_start_session``, etc.) stays synchronous because ``__init__``
runs inside ``asyncio.to_thread``.
Since ``SparkBaseCursor.__init__`` performs I/O (session management),
cursor creation must be wrapped in ``asyncio.to_thread``::
cursor = await asyncio.to_thread(conn.cursor)
Example:
>>> import asyncio
>>> async with await pyathena.aio_connect(
... work_group="spark-workgroup",
... cursor_class=AioSparkCursor,
... ) as conn:
... cursor = await asyncio.to_thread(conn.cursor)
... await cursor.execute("spark.sql('SELECT 1').show()")
... print(await cursor.get_std_out())
"""
@property
@override
def calculation_execution(self) -> AthenaCalculationExecution | None:
return self._calculation_execution
# --- async overrides of SparkBaseCursor I/O methods ---
@override
async def _get_calculation_execution_status( # type: ignore[override]
self, query_id: str
) -> AthenaCalculationExecutionStatus:
request: dict[str, Any] = {"CalculationExecutionId": query_id}
try:
response = await async_retry_api_call(
self._connection.client.get_calculation_execution_status,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to get calculation execution status.")
raise OperationalError(*e.args) from e
else:
return AthenaCalculationExecutionStatus(response)
@override
async def _get_calculation_execution( # type: ignore[override]
self, query_id: str
) -> AthenaCalculationExecution:
request: dict[str, Any] = {"CalculationExecutionId": query_id}
try:
response = await async_retry_api_call(
self._connection.client.get_calculation_execution,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to get calculation execution.")
raise OperationalError(*e.args) from e
else:
return AthenaCalculationExecution(response)
@override
async def _calculate( # type: ignore[override]
self,
session_id: str,
code_block: str,
description: str | None = None,
client_request_token: str | None = None,
) -> str:
"""Start a calculation execution with ``StartCalculationExecution``.
Without ``client_request_token``, a generated token is sent, so that a
retried request returns the calculation an earlier attempt started instead
of starting another one.
With ``kill_on_interrupt`` enabled, the request runs in a task shielded
from task cancellation. On cancellation, the request is abandoned if that
task has not begun it by then; it is never sent, and the cancellation
propagates. Otherwise the cursor waits for the request to finish, requests
cancellation of the calculation it started, waits for a terminal state,
stores the calculation ID and execution on the cursor, and re-raises
``asyncio.CancelledError``. Another cancellation during that wait
propagates at once.
Args:
session_id: The session ID.
code_block: The code to run.
description: The calculation description.
client_request_token: The idempotency token of the request.
Returns:
The calculation execution ID.
Raises:
asyncio.CancelledError: If the task is cancelled while starting the
calculation. A failure to start, cancel, or wait for the
calculation becomes its ``__cause__``.
DatabaseError: If the request fails.
"""
request = self._build_start_calculation_execution_request(
session_id=session_id,
code_block=code_block,
description=description,
client_request_token=client_request_token or str(uuid.uuid4()),
)
if not self._kill_on_interrupt:
return await self._start_calculation_execution(request)
caller = asyncio.current_task()
cancel_requests = caller.cancelling() if caller else 0
async def run() -> str | None:
# Begin the request only if the caller has not been cancelled since.
if caller and caller.cancelling() > cancel_requests:
return None
return await self._start_calculation_execution(request)
start = asyncio.ensure_future(run())
try:
return cast(str, await asyncio.shield(start))
except asyncio.CancelledError as cancellation:
try:
calculation_id = await start
if calculation_id is None:
# The task did not begin the request, so it was never sent.
raise cancellation
_logger.warning("Query canceled by user.")
self._calculation_id = calculation_id
await self._cancel_and_wait(calculation_id)
except Exception as e:
raise cancellation from e
raise
@override
async def _start_calculation_execution( # type: ignore[override]
self, request: dict[str, Any]
) -> str:
"""Send a ``StartCalculationExecution`` request.
Args:
request: The request parameters.
Returns:
The calculation execution ID.
Raises:
DatabaseError: If the request fails.
"""
try:
response = await async_retry_api_call(
self._connection.client.start_calculation_execution,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to execute calculation.")
raise DatabaseError(*e.args) from e
return cast(str, response.get("CalculationExecutionId"))
@override
async def _poll_until_terminal( # type: ignore[override]
self, query_id: str
) -> AthenaQueryExecution | AthenaCalculationExecution:
"""Poll a calculation execution until it reaches a terminal state.
Calls ``on_poll`` with every status and awaits ``poll_interval`` seconds
between requests.
Args:
query_id: The calculation execution ID.
Returns:
The calculation execution in a terminal state.
Raises:
OperationalError: If a status request fails.
"""
while True:
calculation_status = await self._get_calculation_execution_status(query_id)
if self._on_poll:
self._on_poll(calculation_status)
if calculation_status.state in AthenaCalculationExecutionStatus.TERMINAL_STATES:
return await self._get_calculation_execution(query_id)
await asyncio.sleep(self._poll_interval)
@override
async def _poll( # type: ignore[override]
self, query_id: str
) -> AthenaQueryExecution | AthenaCalculationExecution:
"""Wait for a calculation execution to reach a terminal state.
On task cancellation with ``kill_on_interrupt`` enabled, requests
cancellation, waits for the calculation to reach a terminal state, stores
it as the cursor's calculation execution, and re-raises
``asyncio.CancelledError``.
Cancellation is a best-effort request, so the terminal state can be
``COMPLETED`` or ``FAILED`` instead of ``CANCELED``.
Args:
query_id: The calculation execution ID.
Returns:
The calculation execution in a terminal state.
Raises:
asyncio.CancelledError: If the task is cancelled while waiting. A failure
to cancel or wait for the calculation becomes its ``__cause__``.
OperationalError: If a status request fails.
"""
try:
return await self._poll_until_terminal(query_id)
except asyncio.CancelledError as cancellation:
if not self._kill_on_interrupt:
raise
_logger.warning("Query canceled by user.")
try:
await self._cancel_and_wait(query_id)
except Exception as e:
raise cancellation from e
raise
@override
async def _cancel_and_wait(self, calculation_id: str) -> None: # type: ignore[override]
"""Request cancellation and store the calculation's terminal state.
Args:
calculation_id: The calculation execution ID.
Raises:
OperationalError: If the cancellation or a status request fails.
"""
await self._cancel(calculation_id)
self._calculation_execution = cast(
AthenaCalculationExecution, await self._poll_until_terminal(calculation_id)
)
@override
async def _cancel(self, query_id: str) -> None: # type: ignore[override]
request: dict[str, Any] = {"CalculationExecutionId": query_id}
try:
await async_retry_api_call(
self._connection.client.stop_calculation_execution,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception("Failed to cancel calculation.")
raise OperationalError(*e.args) from e
@override
async def _terminate_session(self) -> None: # type: ignore[override]
request: dict[str, Any] = {"SessionId": self._session_id}
try:
await async_retry_api_call(
self._connection.client.terminate_session,
config=self._retry_config,
logger=_logger,
**request,
)
except Exception as e:
_logger.exception(f"Failed to terminate session: {self._session_id}.")
raise OperationalError(*e.args) from e
@override
async def _read_s3_file_as_text(self, uri) -> str: # type: ignore[override]
bucket, key = parse_output_location(uri)
response = await async_retry_api_call(
self._client.get_object,
config=self._retry_config,
logger=_logger,
Bucket=bucket,
Key=key,
)
return cast(str, response["Body"].read().decode("utf-8").strip())
# --- public API ---
[docs]
async def get_std_out(self) -> str | None:
"""Get the standard output from the Spark calculation execution.
Returns:
The standard output as a string, or None if no output is available.
"""
if not self._calculation_execution or not self._calculation_execution.std_out_s3_uri:
return None
return await self._read_s3_file_as_text(self._calculation_execution.std_out_s3_uri)
[docs]
async def get_std_error(self) -> str | None:
"""Get the standard error from the Spark calculation execution.
Returns:
The standard error as a string, or None if no error output is available.
"""
if not self._calculation_execution or not self._calculation_execution.std_error_s3_uri:
return None
return await self._read_s3_file_as_text(self._calculation_execution.std_error_s3_uri)
[docs]
@override
async 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,
) -> AioSparkCursor:
"""Execute PySpark code asynchronously.
Args:
operation: PySpark code to execute.
parameters: Unused, kept for API compatibility.
session_id: Spark session ID override.
description: Calculation description.
client_request_token: Idempotency token.
work_group: Unused, kept for API compatibility.
**kwargs: Additional parameters.
Returns:
Self reference for method chaining.
"""
# A failure below must not leave the previous calculation on the cursor.
self._calculation_id = None
self._calculation_execution = None
self._calculation_id = await 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, await self._poll(self._calculation_id)
)
if self._calculation_execution.state != AthenaCalculationExecutionStatus.STATE_COMPLETED:
std_error = await self.get_std_error()
raise OperationalError(std_error)
return self
[docs]
async def cancel(self) -> None:
"""Cancel the currently running calculation.
Raises:
ProgrammingError: If no calculation is running.
"""
if not self.calculation_id:
raise ProgrammingError("CalculationExecutionId is none or empty.")
await self._cancel(self.calculation_id)
[docs]
@override
async def close(self) -> None: # type: ignore[override]
"""Close the cursor, terminating its Spark session if configured to.
See ``terminate_session_on_close``. After a successful termination,
further calls do not terminate the session again; after a failed one,
calling this method again retries it.
Raises:
OperationalError: If terminating the session fails.
"""
if self._terminate_session_on_close:
await self._terminate_session()
# Terminated; later calls do nothing.
self._terminate_session_on_close = False
[docs]
@override
async def executemany( # type: ignore[override]
self,
operation: str,
seq_of_parameters: list[dict[str, Any] | list[str] | None],
**kwargs,
) -> None:
raise NotSupportedError
def __aiter__(self):
return self
async def __anext__(self):
raise StopAsyncIteration
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
await self.close()