From f4cc848b8502cb4292ca932b8341cf898728454d Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 28 Sep 2026 22:47:14 +0900 Subject: [PATCH] Replace name-mangled helpers with single-underscore methods Rename the 16 non-dunder `__name` helpers so subclasses can reuse and override them, and inline the one-line fetch wrappers: - `__poll` in the SQL and Spark cursors (sync and async) becomes `_poll_until_terminal`. - Spark: `_cancel_and_wait`, `_start_calculation_execution`, `_wait_for_calculation_start`, `_terminate_session_by_id`. - Result sets: `_get_query_results`, `_async_get_query_results`, and `_create_s3_file_system` in the Arrow and pandas result sets. Share the non-I/O parts that the sync and async copies repeated: the terminal states become `TERMINAL_STATES` on the execution models, and `_build_get_query_results_request` holds the GetQueryResults checks and request. The async Spark cursor now logs the session ID when terminating a session fails, as the sync cursor does. Closes #879 Co-Authored-By: Claude Opus 5.5 --- pyathena/aio/common.py | 41 +++++++++++++++++++----- pyathena/aio/result_set.py | 47 ++++++++++++++++++---------- pyathena/aio/spark/cursor.py | 44 +++++++++++++++++--------- pyathena/arrow/result_set.py | 9 ++++-- pyathena/common.py | 43 ++++++++++++++++++++----- pyathena/model.py | 2 ++ pyathena/pandas/result_set.py | 9 ++++-- pyathena/result_set.py | 59 +++++++++++++++++++++++++++++------ pyathena/spark/common.py | 57 ++++++++++++++++++++------------- 9 files changed, 230 insertions(+), 81 deletions(-) diff --git a/pyathena/aio/common.py b/pyathena/aio/common.py index 16547d1e..a2beac17 100644 --- a/pyathena/aio/common.py +++ b/pyathena/aio/common.py @@ -113,27 +113,52 @@ async def _get_query_execution(self, query_id: str) -> AthenaQueryExecution: # else: return AthenaQueryExecution(response) - async def __poll(self, query_id: str) -> AthenaQueryExecution: + async def _poll_until_terminal(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override] + """Poll a query execution until it reaches a terminal state. + + Calls ``on_poll`` with every status and awaits ``poll_interval`` seconds + between requests. + + Args: + query_id: The query execution ID. + + Returns: + The query execution in a terminal state. + + Raises: + OperationalError: If a status request fails. + """ while True: query_execution = await self._get_query_execution(query_id) if self._on_poll: self._on_poll(query_execution) - if query_execution.state in [ - AthenaQueryExecution.STATE_SUCCEEDED, - AthenaQueryExecution.STATE_FAILED, - AthenaQueryExecution.STATE_CANCELLED, - ]: + if query_execution.state in AthenaQueryExecution.TERMINAL_STATES: return query_execution await asyncio.sleep(self._poll_interval) async def _poll(self, query_id: str) -> AthenaQueryExecution: # type: ignore[override] + """Wait for a query execution to finish. + + On ``asyncio.CancelledError`` with ``kill_on_interrupt`` enabled, stops the + query and returns its final execution instead of re-raising. + + Args: + query_id: The query execution ID. + + Returns: + The query execution in a terminal state. + + Raises: + asyncio.CancelledError: If cancelled and ``kill_on_interrupt`` is disabled. + OperationalError: If a status or stop request fails. + """ try: - query_execution = await self.__poll(query_id) + query_execution = await self._poll_until_terminal(query_id) except asyncio.CancelledError: if self._kill_on_interrupt: _logger.warning("Query canceled by user.") await self._cancel(query_id) - query_execution = await self.__poll(query_id) + query_execution = await self._poll_until_terminal(query_id) else: raise return query_execution diff --git a/pyathena/aio/result_set.py b/pyathena/aio/result_set.py index 47fe6508..0f393ff9 100644 --- a/pyathena/aio/result_set.py +++ b/pyathena/aio/result_set.py @@ -85,21 +85,26 @@ async def create( await result_set._async_pre_fetch() return result_set - async def __async_get_query_results( + async def _async_get_query_results( self, max_results: int, next_token: str | None = None ) -> dict[str, Any]: - if not self.query_id: - raise ProgrammingError("QueryExecutionId is none or empty.") - if self.state != AthenaQueryExecution.STATE_SUCCEEDED: - raise ProgrammingError("QueryExecutionState is not SUCCEEDED.") + """Get a page of query results with ``GetQueryResults``. + + Args: + max_results: The maximum number of rows in the page. + next_token: The token of the page to get; the first page if None. + + Returns: + The ``GetQueryResults`` response. + + Raises: + ProgrammingError: If the query ID is missing, the query has not + succeeded, or the result set is closed. + OperationalError: If the request fails. + """ + request = self._build_get_query_results_request(max_results, next_token) if self.is_closed: raise ProgrammingError("AthenaAioResultSet is closed.") - request: dict[str, Any] = { - "QueryExecutionId": self.query_id, - "MaxResults": max_results, - } - if next_token: - request["NextToken"] = next_token try: response = await async_retry_api_call( self.connection.client.get_query_results, @@ -113,18 +118,28 @@ async def __async_get_query_results( else: return cast(dict[str, Any], response) - async def __async_fetch(self, next_token: str | None = None) -> dict[str, Any]: - return await self.__async_get_query_results(self._arraysize, next_token) - async def _async_fetch(self) -> None: + """Fetch the next page of rows into the result set. + + Raises: + ProgrammingError: If there is no next page. + OperationalError: If the request fails. + """ if not self._next_token: raise ProgrammingError("NextToken is none or empty.") - response = await self.__async_fetch(self._next_token) + response = await self._async_get_query_results(self._arraysize, self._next_token) rows, self._next_token = self._parse_result_rows(response) self._process_rows(rows) async def _async_pre_fetch(self) -> None: - response = await self.__async_fetch() + """Fetch the first page of rows along with the result metadata. + + Raises: + ProgrammingError: If the query ID is missing, the query has not + succeeded, or the result set is closed. + OperationalError: If the request fails. + """ + response = await self._async_get_query_results(self._arraysize) self._process_metadata(response) self._process_update_count(response) rows, self._next_token = self._parse_result_rows(response) diff --git a/pyathena/aio/spark/cursor.py b/pyathena/aio/spark/cursor.py index 71cebdd0..de4071fb 100644 --- a/pyathena/aio/spark/cursor.py +++ b/pyathena/aio/spark/cursor.py @@ -169,21 +169,23 @@ async def _calculate( # type: ignore[override] client_request_token=client_request_token or str(uuid.uuid4()), ) if not self._kill_on_interrupt: - return await self.__start_calculation(request) + return await self._start_calculation_execution(request) - start = asyncio.ensure_future(self.__start_calculation(request)) + start = asyncio.ensure_future(self._start_calculation_execution(request)) try: return await asyncio.shield(start) except asyncio.CancelledError as cancellation: _logger.warning("Query canceled by user.") try: self._calculation_id = await start - await self.__cancel_and_wait(self._calculation_id) + await self._cancel_and_wait(self._calculation_id) except Exception as e: raise cancellation from e raise - async def __start_calculation(self, request: dict[str, Any]) -> str: + async def _start_calculation_execution( # type: ignore[override] + self, request: dict[str, Any] + ) -> str: """Send a ``StartCalculationExecution`` request. Args: @@ -207,16 +209,28 @@ async def __start_calculation(self, request: dict[str, Any]) -> str: raise DatabaseError(*e.args) from e return cast(str, response.get("CalculationExecutionId")) - async def __poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: + 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.STATE_COMPLETED, - AthenaCalculationExecutionStatus.STATE_FAILED, - AthenaCalculationExecutionStatus.STATE_CANCELED, - ]: + if calculation_status.state in AthenaCalculationExecutionStatus.TERMINAL_STATES: return await self._get_calculation_execution(query_id) await asyncio.sleep(self._poll_interval) @@ -244,18 +258,18 @@ async def _poll( # type: ignore[override] OperationalError: If a status request fails. """ try: - return await self.__poll(query_id) + 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) + await self._cancel_and_wait(query_id) except Exception as e: raise cancellation from e raise - async def __cancel_and_wait(self, calculation_id: str) -> None: + async def _cancel_and_wait(self, calculation_id: str) -> None: # type: ignore[override] """Request cancellation and store the calculation's terminal state. Args: @@ -266,7 +280,7 @@ async def __cancel_and_wait(self, calculation_id: str) -> None: """ await self._cancel(calculation_id) self._calculation_execution = cast( - AthenaCalculationExecution, await self.__poll(calculation_id) + AthenaCalculationExecution, await self._poll_until_terminal(calculation_id) ) async def _cancel(self, query_id: str) -> None: # type: ignore[override] @@ -292,7 +306,7 @@ async def _terminate_session(self) -> None: # type: ignore[override] **request, ) except Exception as e: - _logger.exception("Failed to terminate session.") + _logger.exception(f"Failed to terminate session: {self._session_id}.") raise OperationalError(*e.args) from e async def _read_s3_file_as_text(self, uri) -> str: # type: ignore[override] diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index 2ff47683..e8cd1ece 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -110,7 +110,7 @@ def __init__( self._connect_timeout = connect_timeout self._request_timeout = request_timeout self._kwargs = kwargs - self._fs = self.__s3_file_system() + self._fs = self._create_s3_file_system() if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: self._table = self._as_arrow() elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: @@ -121,7 +121,12 @@ def __init__( self._table = pa.Table.from_pydict({}) self._batches = iter(self._table.to_batches(arraysize)) - def __s3_file_system(self): + def _create_s3_file_system(self): + """Create a pyarrow ``S3FileSystem`` from the connection settings. + + Returns: + The pyarrow S3 filesystem for reading the query results. + """ from pyarrow import fs connection = self.connection diff --git a/pyathena/common.py b/pyathena/common.py index cfb7d09f..b18868c3 100644 --- a/pyathena/common.py +++ b/pyathena/common.py @@ -862,27 +862,54 @@ def _list_query_executions( return next_token, [] return next_token, self._batch_get_query_execution(query_ids) - def __poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: + def _poll_until_terminal( + self, query_id: str + ) -> AthenaQueryExecution | AthenaCalculationExecution: + """Poll a query execution until it reaches a terminal state. + + Calls ``on_poll`` with every status and sleeps ``poll_interval`` seconds + between requests. + + Args: + query_id: The query execution ID. + + Returns: + The query execution in a terminal state. + + Raises: + OperationalError: If a status request fails. + """ while True: query_execution = self._get_query_execution(query_id) if self._on_poll: self._on_poll(query_execution) - if query_execution.state in [ - AthenaQueryExecution.STATE_SUCCEEDED, - AthenaQueryExecution.STATE_FAILED, - AthenaQueryExecution.STATE_CANCELLED, - ]: + if query_execution.state in AthenaQueryExecution.TERMINAL_STATES: return query_execution time.sleep(self._poll_interval) def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: + """Wait for a query execution to finish. + + On ``KeyboardInterrupt`` with ``kill_on_interrupt`` enabled, stops the query + and returns its final execution instead of re-raising. + + Args: + query_id: The query execution ID. + + Returns: + The query execution in a terminal state. + + Raises: + KeyboardInterrupt: If interrupted and ``kill_on_interrupt`` is disabled. + OperationalError: If a status or stop request fails. + """ try: - query_execution = self.__poll(query_id) + query_execution = self._poll_until_terminal(query_id) except KeyboardInterrupt as e: if self._kill_on_interrupt: _logger.warning("Query canceled by user.") self._cancel(query_id) - query_execution = self.__poll(query_id) + query_execution = self._poll_until_terminal(query_id) else: raise e return query_execution diff --git a/pyathena/model.py b/pyathena/model.py index f170622f..6cbb335d 100644 --- a/pyathena/model.py +++ b/pyathena/model.py @@ -49,6 +49,7 @@ class AthenaQueryExecution: STATE_SUCCEEDED: str = "SUCCEEDED" STATE_FAILED: str = "FAILED" STATE_CANCELLED: str = "CANCELLED" + TERMINAL_STATES: tuple[str, ...] = (STATE_SUCCEEDED, STATE_FAILED, STATE_CANCELLED) STATEMENT_TYPE_DDL: str = "DDL" STATEMENT_TYPE_DML: str = "DML" @@ -353,6 +354,7 @@ class AthenaCalculationExecutionStatus: STATE_CANCELED: str = "CANCELED" STATE_COMPLETED: str = "COMPLETED" STATE_FAILED: str = "FAILED" + TERMINAL_STATES: tuple[str, ...] = (STATE_COMPLETED, STATE_FAILED, STATE_CANCELED) def __init__(self, response: dict[str, Any]) -> None: status = response.get("Status") diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 4eba6d60..b297cd5a 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -304,7 +304,7 @@ def __init__( self._auto_optimize_chunksize = auto_optimize_chunksize self._data_manifest: list[str] = [] self._kwargs = kwargs - self._fs = self.__s3_file_system() + self._fs = self._create_s3_file_system() self._csv_stream: TextIOWrapper | None = None # Cache time column names for efficient _trunc_date processing @@ -427,7 +427,12 @@ def _auto_determine_chunksize(self, file_size_bytes: int) -> int | None: return self.AUTO_CHUNK_SIZE_MEDIUM return None - def __s3_file_system(self): + def _create_s3_file_system(self): + """Create PyAthena's ``S3FileSystem`` from the connection settings. + + Returns: + The S3 filesystem for reading the query results. + """ from pyathena.filesystem.s3 import S3FileSystem return S3FileSystem( diff --git a/pyathena/result_set.py b/pyathena/result_set.py index f6fe6017..d75326c3 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -343,21 +343,52 @@ def connection(self) -> Connection[Any]: raise ProgrammingError("AthenaResultSet is closed.") return cast("Connection[Any]", self._connection) - def __get_query_results( + def _build_get_query_results_request( self, max_results: int, next_token: str | None = None ) -> dict[str, Any]: + """Build a ``GetQueryResults`` request for the result set's query. + + Args: + max_results: The maximum number of rows in the page. + next_token: The token of the page to get; the first page if None. + + Returns: + The request parameters. + + Raises: + ProgrammingError: If the query ID is missing or the query has not + succeeded. + """ if not self.query_id: raise ProgrammingError("QueryExecutionId is none or empty.") if self.state != AthenaQueryExecution.STATE_SUCCEEDED: raise ProgrammingError("QueryExecutionState is not SUCCEEDED.") - if self.is_closed: - raise ProgrammingError("AthenaResultSet is closed.") request: dict[str, Any] = { "QueryExecutionId": self.query_id, "MaxResults": max_results, } if next_token: request["NextToken"] = next_token + return request + + def _get_query_results(self, max_results: int, next_token: str | None = None) -> dict[str, Any]: + """Get a page of query results with ``GetQueryResults``. + + Args: + max_results: The maximum number of rows in the page. + next_token: The token of the page to get; the first page if None. + + Returns: + The ``GetQueryResults`` response. + + Raises: + ProgrammingError: If the query ID is missing, the query has not + succeeded, or the result set is closed. + OperationalError: If the request fails. + """ + request = self._build_get_query_results_request(max_results, next_token) + if self.is_closed: + raise ProgrammingError("AthenaResultSet is closed.") try: response = retry_api_call( self.connection.client.get_query_results, @@ -371,18 +402,28 @@ def __get_query_results( else: return cast(dict[str, Any], response) - def __fetch(self, next_token: str | None = None) -> dict[str, Any]: - return self.__get_query_results(self._arraysize, next_token) - def _fetch(self) -> None: + """Fetch the next page of rows into the result set. + + Raises: + ProgrammingError: If there is no next page. + OperationalError: If the request fails. + """ if not self._next_token: raise ProgrammingError("NextToken is none or empty.") - response = self.__fetch(self._next_token) + response = self._get_query_results(self._arraysize, self._next_token) rows, self._next_token = self._parse_result_rows(response) self._process_rows(rows) def _pre_fetch(self) -> None: - response = self.__fetch() + """Fetch the first page of rows along with the result metadata. + + Raises: + ProgrammingError: If the query ID is missing, the query has not + succeeded, or the result set is closed. + OperationalError: If the request fails. + """ + response = self._get_query_results(self._arraysize) self._process_metadata(response) self._process_update_count(response) rows, self._next_token = self._parse_result_rows(response) @@ -600,7 +641,7 @@ def _fetch_all_rows( next_token: str | None = None while True: - response = self.__get_query_results(self.DEFAULT_FETCH_SIZE, next_token) + response = self._get_query_results(self.DEFAULT_FETCH_SIZE, next_token) rows, next_token = self._parse_result_rows(response) offset = 1 if rows and self._is_first_row_column_labels(rows) else 0 diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 8114dc9b..01637313 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -272,7 +272,7 @@ def _start_session(self) -> str: # The caller receives no cursor to close, so the session is released here. with contextlib.suppress(OperationalError): # Already logged with the session ID; the original error takes precedence. - self.__terminate_session(session_id) + self._terminate_session_by_id(session_id) raise return session_id @@ -282,13 +282,14 @@ def _terminate_session(self) -> None: Raises: OperationalError: If the request fails. """ - self.__terminate_session(self._session_id) + self._terminate_session_by_id(self._session_id) - def __terminate_session(self, session_id: str) -> None: + def _terminate_session_by_id(self, session_id: str) -> None: """Terminate a Spark session with ``TerminateSession``. Session startup calls this synchronously in every cursor variant, - including those that override ``_terminate_session`` with a coroutine. + including those that override ``_terminate_session`` with a coroutine, + so subclasses must not override this method with a coroutine. Args: session_id: The session ID. @@ -308,16 +309,28 @@ def __terminate_session(self, session_id: str) -> None: _logger.exception(f"Failed to terminate session: {session_id}.") raise OperationalError(*e.args) from e - def __poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: + def _poll_until_terminal( + self, query_id: str + ) -> AthenaQueryExecution | AthenaCalculationExecution: + """Poll a calculation execution until it reaches a terminal state. + + Calls ``on_poll`` with every status and sleeps ``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 = self._get_calculation_execution_status(query_id) if self._on_poll: self._on_poll(calculation_status) - if calculation_status.state in [ - AthenaCalculationExecutionStatus.STATE_COMPLETED, - AthenaCalculationExecutionStatus.STATE_FAILED, - AthenaCalculationExecutionStatus.STATE_CANCELED, - ]: + if calculation_status.state in AthenaCalculationExecutionStatus.TERMINAL_STATES: return self._get_calculation_execution(query_id) time.sleep(self._poll_interval) @@ -342,18 +355,18 @@ def _poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecut OperationalError: If a status request fails. """ try: - return self.__poll(query_id) + return self._poll_until_terminal(query_id) except KeyboardInterrupt as interrupt: if not self._kill_on_interrupt: raise _logger.warning("Query canceled by user.") try: - self.__cancel_and_wait(query_id) + self._cancel_and_wait(query_id) except Exception as e: raise interrupt from e raise - def __cancel_and_wait(self, calculation_id: str) -> None: + def _cancel_and_wait(self, calculation_id: str) -> None: """Request cancellation and store the calculation's terminal state. Args: @@ -363,7 +376,9 @@ def __cancel_and_wait(self, calculation_id: str) -> None: OperationalError: If the cancellation or a status request fails. """ self._cancel(calculation_id) - self._calculation_execution = cast(AthenaCalculationExecution, self.__poll(calculation_id)) + self._calculation_execution = cast( + AthenaCalculationExecution, self._poll_until_terminal(calculation_id) + ) def _calculate( self, @@ -409,7 +424,7 @@ def _calculate( client_request_token=client_request_token or str(uuid.uuid4()), ) if not self._kill_on_interrupt: - return self.__start_calculation(request) + return self._start_calculation_execution(request) future: Future[str] = Future() @@ -418,26 +433,26 @@ def start() -> None: if not future.set_running_or_notify_cancel(): return try: - future.set_result(self.__start_calculation(request)) + future.set_result(self._start_calculation_execution(request)) except BaseException as e: future.set_exception(e) try: threading.Thread(target=start, name="pyathena-spark-start", daemon=True).start() - return self.__wait_for_start(future) + return self._wait_for_calculation_start(future) except KeyboardInterrupt as interrupt: if future.cancel(): # The helper has not begun the request and never will. raise _logger.warning("Query canceled by user.") try: - self._calculation_id = self.__wait_for_start(future) - self.__cancel_and_wait(self._calculation_id) + self._calculation_id = self._wait_for_calculation_start(future) + self._cancel_and_wait(self._calculation_id) except Exception as e: raise interrupt from e raise - def __start_calculation(self, request: dict[str, Any]) -> str: + def _start_calculation_execution(self, request: dict[str, Any]) -> str: """Send a ``StartCalculationExecution`` request. Args: @@ -462,7 +477,7 @@ def __start_calculation(self, request: dict[str, Any]) -> str: return cast(str, response.get("CalculationExecutionId")) @staticmethod - def __wait_for_start(future: Future[str]) -> str: + def _wait_for_calculation_start(future: Future[str]) -> str: """Wait for the start request on a helper thread to finish. Args: