From 899d1a52f5da72474bcb52a4af19703944133014 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 20:59:11 +0900 Subject: [PATCH 1/7] Backport the CI test policy and release gate to 3.x Port the policies of #789, #837, #862, and #863 to the 3.x workflows without the parallel suites or the per-change suite selection: - A lint job runs for every pull request, including Draft and fork ones. - Draft and external-fork pull requests run no AWS suites; a ready pull request runs them on the newest Python version only. - workflow_dispatch accepts a python-versions input and otherwise runs every supported version. - The Release workflow runs every suite on every supported Python version for the tagged commit through workflow_call before building and publishing, so a failing tag publishes nothing. The fixed S3 Tables namespace no longer exists in the test account, so AWS_ATHENA_S3_TABLES_NAMESPACE is dropped and the SQLAlchemy S3 Tables tests skip on this branch. The weekly schedule is removed because scheduled runs only use the default branch's workflow. Co-Authored-By: Claude Opus 5.5 --- .github/workflows/release.yaml | 18 +++++-- .github/workflows/test-suite.yaml | 15 ++++-- .github/workflows/test.yaml | 86 ++++++++++++++++++++++++++++--- 3 files changed, 103 insertions(+), 16 deletions(-) diff --git a/.github/workflows/release.yaml b/.github/workflows/release.yaml index d79188242..0ca9581a9 100644 --- a/.github/workflows/release.yaml +++ b/.github/workflows/release.yaml @@ -5,13 +5,23 @@ on: tags: - 'v*' -permissions: - id-token: write - contents: write +permissions: {} jobs: + # Runs every suite on every supported Python version for the tagged commit; + # nothing is built or published unless all of them pass. + test: + uses: ./.github/workflows/test.yaml + permissions: + contents: read + id-token: write + release: + needs: test runs-on: ubuntu-latest + permissions: + id-token: write + contents: write env: PYTHON_VERSION: '3.12' @@ -22,7 +32,7 @@ jobs: - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0 with: - python-version: ${{ matrix.python-version }} + python-version: ${{ env.PYTHON_VERSION }} enable-cache: true - name: Build diff --git a/.github/workflows/test-suite.yaml b/.github/workflows/test-suite.yaml index 535bf9a80..22be768a9 100644 --- a/.github/workflows/test-suite.yaml +++ b/.github/workflows/test-suite.yaml @@ -6,9 +6,15 @@ on: test-type: required: true type: string + python-versions: + description: JSON array of the Python versions to test + required: true + type: string jobs: run: + # External fork contributions must validate AWS behavior in their own account. + if: github.event_name != 'pull_request' || github.event.pull_request.head.repo.full_name == github.repository runs-on: ubuntu-latest env: @@ -18,16 +24,15 @@ jobs: AWS_ATHENA_WORKGROUP: pyathena AWS_ATHENA_SPARK_WORKGROUP: pyathena-spark AWS_ATHENA_MANAGED_WORKGROUP: pyathena-managed - # Registered S3 Tables catalog (s3tablescatalog/) and namespace - # for the SQLAlchemy S3 Tables tests; the table bucket, namespace, and the - # AWS analytics-services integration are provisioned out of band. + # The SQLAlchemy S3 Tables tests need a fixed namespace, which the test + # account no longer provides (master creates one per test session), so + # they are skipped on this branch. AWS_ATHENA_S3_TABLES_CATALOG: s3tablescatalog/laughingman7743-pyathena-s3-tables - AWS_ATHENA_S3_TABLES_NAMESPACE: pyathena strategy: fail-fast: false matrix: - python-version: ['3.10', '3.11', '3.12', '3.13', '3.14'] + python-version: ${{ fromJSON(inputs.python-versions) }} steps: - name: Checkout diff --git a/.github/workflows/test.yaml b/.github/workflows/test.yaml index d8e6a3eec..1ab76af24 100644 --- a/.github/workflows/test.yaml +++ b/.github/workflows/test.yaml @@ -2,15 +2,23 @@ name: Test on: pull_request: + # ready_for_review starts the AWS jobs for a pull request leaving Draft; + # converted_to_draft starts a run without them, which cancels an + # in-progress run through the concurrency group. + types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] paths-ignore: - 'docs/**' - '**.md' - schedule: - - cron: '0 0 * * 0' - # Allows refreshing the README status badge on demand: the badge reflects - # the latest run on the default branch, which is otherwise only the weekly - # scheduled run and stays red for up to a week after a transient failure. + # Runs every suite on the selected branch, on the requested Python versions. workflow_dispatch: + inputs: + python-versions: + description: Comma-separated Python versions, such as 3.12 or 3.11,3.14; empty for every supported version + type: string + default: '' + # The Release workflow runs every suite on every supported Python version + # before publishing. + workflow_call: permissions: id-token: write @@ -23,19 +31,83 @@ concurrency: cancel-in-progress: true jobs: + # Offline checks run for every event, including Draft and fork pull requests. + lint: + runs-on: ubuntu-latest + permissions: + contents: read + steps: + - name: Checkout + uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + persist-credentials: false + - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0 + with: + python-version: '3.12' + enable-cache: true + - uses: taiki-e/install-action@7a79fe8c3a13344501c80d99cae481c1c9085912 # v2.81.10 + with: + tool: just + - run: just lint + + # Selects the Python versions of the AWS suites. Draft and external-fork pull + # requests run none. A ready pull request tests the newest Python version; a + # dispatch tests the requested versions or every version, and the Release + # workflow every version. + versions: + if: >- + github.event_name != 'pull_request' || + (!github.event.pull_request.draft && + github.event.pull_request.head.repo.full_name == github.repository) + runs-on: ubuntu-latest + permissions: {} + outputs: + python-versions: ${{ steps.select.outputs.python-versions }} + steps: + - id: select + env: + EVENT_NAME: ${{ github.event_name }} + REQUESTED_VERSIONS: ${{ inputs.python-versions }} + # Every supported version, oldest first; keep in sync with the + # pyproject.toml classifiers. + PYTHON_VERSIONS: '["3.10", "3.11", "3.12", "3.13", "3.14"]' + run: | + case "$EVENT_NAME" in + pull_request) + versions=$(jq -c '[last]' <<< "$PYTHON_VERSIONS") + ;; + workflow_dispatch) + versions=$(jq -c --arg requested "$REQUESTED_VERSIONS" ' + ($requested | split(",") | map(gsub("\\s"; "")) | map(select(. != "")) | unique) as $selected + | if $selected == [] then . + elif ($selected - .) == [] then $selected + else error("unsupported Python versions: \($selected - . | join(", "))") + end' <<< "$PYTHON_VERSIONS") + ;; + *) + # The Release workflow (a workflow_call from a tag push). + versions=$(jq -c '.' <<< "$PYTHON_VERSIONS") + ;; + esac + echo "python-versions=$versions" >> "$GITHUB_OUTPUT" + test: + needs: versions uses: ./.github/workflows/test-suite.yaml with: test-type: pyathena + python-versions: ${{ needs.versions.outputs.python-versions }} test-sqla: - needs: [test] + needs: [versions, test] uses: ./.github/workflows/test-suite.yaml with: test-type: sqla + python-versions: ${{ needs.versions.outputs.python-versions }} test-sqla-async: - needs: [test-sqla] + needs: [versions, test-sqla] uses: ./.github/workflows/test-suite.yaml with: test-type: sqla_async + python-versions: ${{ needs.versions.outputs.python-versions }} From e6b2f9f3c301231e4579b6f908e393db932ead74 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 20:45:13 +0900 Subject: [PATCH 2/7] Backport #800: Find an existing table in to_sql whatever the name's case (cherry picked from commit ad1232b2eb19c67c8793a3f86ea7cbdd52366a47) Co-Authored-By: Claude Opus 5.5 --- pyathena/pandas/util.py | 8 ++++++-- tests/pyathena/pandas/test_util.py | 4 +++- 2 files changed, 9 insertions(+), 3 deletions(-) diff --git a/pyathena/pandas/util.py b/pyathena/pandas/util.py index b2b96612d..9337ce326 100644 --- a/pyathena/pandas/util.py +++ b/pyathena/pandas/util.py @@ -261,13 +261,17 @@ def to_sql( ).Bucket(bucket_name) cursor = conn.cursor() + # Athena stores identifiers in lowercase and information_schema reports them + # that way, so compare lowercase literals. + schema_literal = schema.lower().replace("'", "''") + name_literal = name.lower().replace("'", "''") table = cursor.execute( textwrap.dedent( f""" SELECT table_name FROM information_schema.tables - WHERE table_schema = '{schema}' - AND table_name = '{name}' + WHERE table_schema = '{schema_literal}' + AND table_name = '{name_literal}' """ ) ).fetchall() diff --git a/tests/pyathena/pandas/test_util.py b/tests/pyathena/pandas/test_util.py index 788308da6..1acbc07ec 100644 --- a/tests/pyathena/pandas/test_util.py +++ b/tests/pyathena/pandas/test_util.py @@ -342,7 +342,9 @@ def test_to_sql(cursor): "col_binary", ] ] - table_name = f"""to_sql_{str(uuid.uuid4()).replace("-", "")}""" + # Uppercase letters: Athena reports names in lowercase, and the existence + # check behind if_exists has to find the table anyway. + table_name = f"To_Sql_{uuid.uuid4().hex}" location = f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/" to_sql( df, From de40c6f8e5c944f0e9e96f124b905a90716b22aa Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 20:45:14 +0900 Subject: [PATCH 3/7] Backport #830: Shut down the AsyncSparkCursor executor when session termination fails (cherry picked from commit 2a107adf1f68713e7b7da2e334cc98ca543b7e7d) Co-Authored-By: Claude Opus 5.5 --- pyathena/spark/async_cursor.py | 18 ++++- tests/pyathena/spark/test_async_cursor.py | 94 +++++++++++++++++++++++ 2 files changed, 110 insertions(+), 2 deletions(-) diff --git a/pyathena/spark/async_cursor.py b/pyathena/spark/async_cursor.py index 1a476f208..b141cd1df 100644 --- a/pyathena/spark/async_cursor.py +++ b/pyathena/spark/async_cursor.py @@ -87,8 +87,22 @@ def __init__( self._executor = ThreadPoolExecutor(max_workers=max_workers) def close(self, wait: bool = False) -> None: - super().close() - self._executor.shutdown(wait=wait) + """Terminate the Spark session, then shut down the executor. + + The executor is shut down even if terminating the session fails. + If termination fails, calling this method again retries it. + + 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) def calculation_execution(self, query_id: str) -> "Future[AthenaCalculationExecution]": return self._executor.submit(self._get_calculation_execution, query_id) diff --git a/tests/pyathena/spark/test_async_cursor.py b/tests/pyathena/spark/test_async_cursor.py index e921bdc3e..6f9f2d25f 100644 --- a/tests/pyathena/spark/test_async_cursor.py +++ b/tests/pyathena/spark/test_async_cursor.py @@ -1,10 +1,20 @@ import textwrap +import threading import time +from concurrent.futures import ThreadPoolExecutor from random import randint +from unittest.mock import MagicMock +import pytest + +from pyathena import OperationalError from pyathena.model import AthenaCalculationExecutionStatus +from pyathena.spark.async_cursor import AsyncSparkCursor from tests import ENV +# Bounds how long the executor test task blocks when nothing releases it. +_TIMEOUT = 10 + class TestAsyncSparkCursor: def test_spark_dataframe(self, async_spark_cursor): @@ -121,3 +131,87 @@ def test_cancel(self, async_spark_cursor): calculation_execution = future.result() assert calculation_execution.state == AthenaCalculationExecutionStatus.STATE_CANCELED + + @staticmethod + def _cursor_with_submitted_work(): + """Build a cursor whose executor holds one running and one queued future. + + Returns: + A tuple of the cursor, the event that releases the running future, + the running future, and the queued future. + """ + cursor = AsyncSparkCursor.__new__(AsyncSparkCursor) # bypass __init__ to avoid AWS calls + cursor._executor = MagicMock(wraps=ThreadPoolExecutor(max_workers=1)) + started = threading.Event() + release = threading.Event() + + def run(): + started.set() + release.wait(_TIMEOUT) + return "running" + + running = cursor._executor.submit(run) + queued = cursor._executor.submit(lambda: "queued") + assert started.wait(_TIMEOUT) + return cursor, release, running, queued + + @pytest.mark.parametrize("fails", [False, True]) + def test_close_wait_shuts_down_executor(self, fails): + cursor, release, running, queued = self._cursor_with_submitted_work() + + def terminate(): + # Let the running future finish so that shutdown(wait=True) returns. + release.set() + if fails: + raise OperationalError("termination failed") + + cursor._terminate_session = MagicMock(side_effect=terminate) + if fails: + with pytest.raises(OperationalError, match="termination failed"): + cursor.close(wait=True) + else: + cursor.close(wait=True) + + cursor._terminate_session.assert_called_once_with() + cursor._executor.shutdown.assert_called_once_with(wait=True) + assert running.done() + assert running.result() == "running" + assert queued.done() + assert queued.result() == "queued" + with pytest.raises(RuntimeError): + cursor._executor.submit(lambda: None) + + @pytest.mark.parametrize("fails", [False, True]) + def test_close_no_wait_shuts_down_executor(self, fails): + cursor, release, running, queued = self._cursor_with_submitted_work() + cursor._terminate_session = MagicMock( + side_effect=OperationalError("termination failed") if fails else None + ) + try: + if fails: + with pytest.raises(OperationalError, match="termination failed"): + cursor.close(wait=False) + else: + cursor.close(wait=False) + + cursor._terminate_session.assert_called_once_with() + cursor._executor.shutdown.assert_called_once_with(wait=False) + with pytest.raises(RuntimeError): + cursor._executor.submit(lambda: None) + finally: + release.set() + assert running.result(_TIMEOUT) == "running" + assert queued.result(_TIMEOUT) == "queued" + + def test_close_retries_termination_after_failure(self): + cursor = AsyncSparkCursor.__new__(AsyncSparkCursor) # bypass __init__ to avoid AWS calls + cursor._terminate_session = MagicMock( + side_effect=[OperationalError("termination failed"), None] + ) + cursor._executor = ThreadPoolExecutor(max_workers=1) + + with pytest.raises(OperationalError, match="termination failed"): + cursor.close() + cursor.close() + + assert cursor._terminate_session.call_count == 2 From d874b0785f583d4cbd677dccbd35cd391e9de378 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 20:45:59 +0900 Subject: [PATCH 4/7] Backport #831: Stop Spark session readiness polling on failure states (cherry picked from commit 37999c3c6f52365a25297782d544d8412bab4801) Co-Authored-By: Claude Opus 5.5 --- pyathena/spark/common.py | 35 ++++++++-- tests/pyathena/spark/test_common.py | 99 +++++++++++++++++++++++++++++ 2 files changed, 130 insertions(+), 4 deletions(-) create mode 100644 tests/pyathena/spark/test_common.py diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index d25c0b0f2..56b0c3541 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -126,21 +126,48 @@ def _get_session_status(self, session_id: str): else: return AthenaSessionStatus(response) - def _wait_for_idle_session(self, session_id: str): + def _wait_for_idle_session(self, session_id: str) -> None: + """Poll a Spark session with ``GetSessionStatus`` until it is idle. + + Args: + session_id: The session ID. + + Raises: + OperationalError: If the session is terminated, degraded, or failed, + or if the request fails. + """ while True: session_status = self._get_session_status(session_id) if session_status.state in [AthenaSessionStatus.STATE_IDLE]: break - if session_status in [ + if session_status.state in [ AthenaSessionStatus.STATE_TERMINATED, AthenaSessionStatus.STATE_DEGRADED, AthenaSessionStatus.STATE_FAILED, ]: - raise OperationalError(session_status.state_change_reason) + message = f"Session: {session_id} is {session_status.state}." + if session_status.state_change_reason: + message += f" {session_status.state_change_reason}" + raise OperationalError(message) time.sleep(self._poll_interval) def _exists_session(self, session_id: str) -> bool: - request = {"SessionId": session_id} + """Whether a Spark session exists, from ``GetSession``. + + Waits for an existing session to become idle before returning. + + Args: + session_id: The session ID. + + Returns: + True if the session exists; False if Athena rejects it with + ``InvalidRequestException``. + + Raises: + OperationalError: If the request fails for another reason, or if the + session is terminated, degraded, or failed. + """ + request: dict[str, Any] = {"SessionId": session_id} try: retry_api_call( self._connection.client.get_session, diff --git a/tests/pyathena/spark/test_common.py b/tests/pyathena/spark/test_common.py new file mode 100644 index 000000000..f33326d7e --- /dev/null +++ b/tests/pyathena/spark/test_common.py @@ -0,0 +1,99 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT + +from unittest.mock import MagicMock, patch + +import pytest + +from pyathena import OperationalError +from pyathena.model import AthenaSessionStatus +from pyathena.spark.cursor import SparkCursor +from pyathena.util import RetryConfig + + +def _session_status(state: str, reason: str | None = None) -> AthenaSessionStatus: + return AthenaSessionStatus( + {"SessionId": "session_id", "Status": {"State": state, "StateChangeReason": reason}} + ) + + +def _cursor() -> SparkCursor: + cursor = SparkCursor.__new__(SparkCursor) # bypass __init__ to avoid AWS calls + cursor._connection = MagicMock() + cursor._retry_config = RetryConfig() + cursor._poll_interval = 0 + return cursor + + +class TestSparkBaseCursor: + @pytest.mark.parametrize( + "state", + [ + AthenaSessionStatus.STATE_TERMINATED, + AthenaSessionStatus.STATE_DEGRADED, + AthenaSessionStatus.STATE_FAILED, + ], + ) + def test_wait_for_idle_session_raises_on_failure_state(self, state): + cursor = _cursor() + with ( + patch.object( + SparkCursor, + "_get_session_status", + return_value=_session_status(state, "session failure reason"), + ), + patch("pyathena.spark.common.time.sleep", side_effect=AssertionError("slept")), + pytest.raises( + OperationalError, + match=rf"^Session: session_id is {state}\. session failure reason$", + ), + ): + cursor._wait_for_idle_session("session_id") + + def test_wait_for_idle_session_raises_without_reason(self): + cursor = _cursor() + with ( + patch.object( + SparkCursor, + "_get_session_status", + return_value=_session_status(AthenaSessionStatus.STATE_TERMINATED), + ), + patch("pyathena.spark.common.time.sleep", side_effect=AssertionError("slept")), + pytest.raises(OperationalError, match=r"^Session: session_id is TERMINATED\.$"), + ): + cursor._wait_for_idle_session("session_id") + + def test_wait_for_idle_session_waits_until_idle(self): + cursor = _cursor() + statuses = [ + _session_status(AthenaSessionStatus.STATE_CREATING), + _session_status(AthenaSessionStatus.STATE_BUSY), + _session_status(AthenaSessionStatus.STATE_IDLE), + ] + with ( + patch.object(SparkCursor, "_get_session_status", side_effect=statuses) as get_status, + patch("pyathena.spark.common.time.sleep") as sleep, + ): + cursor._wait_for_idle_session("session_id") + assert get_status.call_count == 3 + assert sleep.call_count == 2 + + def test_exists_session_raises_on_failure_state(self): + cursor = _cursor() + with ( + patch.object( + SparkCursor, + "_get_session_status", + return_value=_session_status( + AthenaSessionStatus.STATE_TERMINATED, "session failure reason" + ), + ), + patch("pyathena.spark.common.time.sleep", side_effect=AssertionError("slept")), + pytest.raises(OperationalError, match="session failure reason"), + ): + cursor._exists_session("session_id") + cursor._connection.client.get_session.assert_called_once_with(SessionId="session_id") From 1b5b2cbd01c3566d1aa965d76844e4f067ef2d84 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 20:46:20 +0900 Subject: [PATCH 5/7] Backport #832: Terminate a newly started Spark session when cursor setup fails (cherry picked from commit 5766a13ae82a62664fba47f35290be107cbf16bd) The session termination helper uses the single-underscore name _terminate_session_by_id that master adopted in #881. Co-Authored-By: Claude Opus 5.5 --- pyathena/spark/async_cursor.py | 23 ++++- pyathena/spark/common.py | 83 ++++++++++++++--- tests/pyathena/spark/test_common.py | 139 ++++++++++++++++++++++++++++ 3 files changed, 230 insertions(+), 15 deletions(-) diff --git a/pyathena/spark/async_cursor.py b/pyathena/spark/async_cursor.py index b141cd1df..17aa2a3e0 100644 --- a/pyathena/spark/async_cursor.py +++ b/pyathena/spark/async_cursor.py @@ -75,6 +75,27 @@ def __init__( max_workers: int = (cpu_count() or 1) * 5, **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. + **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, @@ -83,8 +104,6 @@ def __init__( session_idle_timeout_minutes=session_idle_timeout_minutes, **kwargs, ) - self._max_workers = max_workers - self._executor = ThreadPoolExecutor(max_workers=max_workers) def close(self, wait: bool = False) -> None: """Terminate the Spark session, then shut down the executor. diff --git a/pyathena/spark/common.py b/pyathena/spark/common.py index 56b0c3541..7ffbb7bb4 100644 --- a/pyathena/spark/common.py +++ b/pyathena/spark/common.py @@ -1,5 +1,6 @@ from __future__ import annotations +import contextlib import logging import time from abc import ABCMeta, abstractmethod @@ -56,6 +57,22 @@ def __init__( session_idle_timeout_minutes: int | None = None, **kwargs, ) -> None: + """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. + Defaults to ``get_default_engine_configuration()``. + notebook_version: Notebook version of a new session. + session_idle_timeout_minutes: Idle timeout of a new session in minutes. + **kwargs: Arguments passed to ``BaseCursor``. + + Raises: + OperationalError: If the supplied session does not exist, or the + session cannot be started or does not become idle. + """ super().__init__(**kwargs) self._engine_configuration = ( engine_configuration @@ -65,18 +82,11 @@ def __init__( self._notebook_version = notebook_version self._session_description = description self._session_idle_timeout_minutes = session_idle_timeout_minutes - - if session_id: - if self._exists_session(session_id): - self._session_id = session_id - else: - raise OperationalError(f"Session: {session_id} not found.") - else: - self._session_id = self._start_session() - self._calculation_id: str | None = None self._calculation_execution: AthenaCalculationExecution | None = None + # Created before the session so that a local failure cannot leave + # a newly started session behind. self._client = self.connection.session.client( "s3", region_name=self.connection.region_name, @@ -84,6 +94,14 @@ def __init__( **self.connection._client_kwargs, ) + if session_id: + if self._exists_session(session_id): + self._session_id = session_id + else: + raise OperationalError(f"Session: {session_id} not found.") + else: + self._session_id = self._start_session() + @property def session_id(self) -> str: return self._session_id @@ -188,6 +206,19 @@ def _exists_session(self, session_id: str) -> bool: return True def _start_session(self) -> str: + """Start a Spark session with ``StartSession`` and wait until it is idle. + + If waiting for the new session raises, including ``KeyboardInterrupt``, + the session is terminated on a best-effort basis before the exception is + re-raised. + + Returns: + The ID of the new session. + + Raises: + OperationalError: If the session cannot be started or does not + become idle. + """ request: dict[str, Any] = { "WorkGroup": self._work_group, "EngineConfiguration": self._engine_configuration, @@ -208,12 +239,38 @@ def _start_session(self) -> str: except Exception as e: _logger.exception("Failed to start session.") raise OperationalError(*e.args) from e - else: + + try: self._wait_for_idle_session(session_id) - return session_id + except BaseException: + # 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_by_id(session_id) + raise + return session_id def _terminate_session(self) -> None: - request = {"SessionId": self._session_id} + """Terminate the cursor's Spark session with ``TerminateSession``. + + Raises: + OperationalError: If the request fails. + """ + self._terminate_session_by_id(self._session_id) + + 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. + + Args: + session_id: The session ID. + + Raises: + OperationalError: If the request fails. + """ + request: dict[str, Any] = {"SessionId": session_id} try: retry_api_call( self._connection.client.terminate_session, @@ -222,7 +279,7 @@ def _terminate_session(self) -> None: **request, ) except Exception as e: - _logger.exception("Failed to terminate session.") + _logger.exception(f"Failed to terminate session: {session_id}.") raise OperationalError(*e.args) from e def __poll(self, query_id: str) -> AthenaQueryExecution | AthenaCalculationExecution: diff --git a/tests/pyathena/spark/test_common.py b/tests/pyathena/spark/test_common.py index f33326d7e..9dbfd27f3 100644 --- a/tests/pyathena/spark/test_common.py +++ b/tests/pyathena/spark/test_common.py @@ -5,15 +5,22 @@ # # SPDX-License-Identifier: MIT +import logging from unittest.mock import MagicMock, patch import pytest +from botocore.exceptions import ClientError from pyathena import OperationalError +from pyathena.aio.spark.cursor import AioSparkCursor from pyathena.model import AthenaSessionStatus +from pyathena.spark.async_cursor import AsyncSparkCursor +from pyathena.spark.common import SparkBaseCursor from pyathena.spark.cursor import SparkCursor from pyathena.util import RetryConfig +SPARK_CURSOR_CLASSES = [SparkCursor, AsyncSparkCursor, AioSparkCursor] + def _session_status(state: str, reason: str | None = None) -> AthenaSessionStatus: return AthenaSessionStatus( @@ -29,6 +36,36 @@ def _cursor() -> SparkCursor: return cursor +def _connection(): + connection = MagicMock() + connection.client.start_session.return_value = {"SessionId": "new-session"} + connection.client.get_session.return_value = {"SessionId": "supplied-session"} + connection.client.get_session_status.return_value = { + "Status": {"State": AthenaSessionStatus.STATE_IDLE} + } + return connection + + +def _init_cursor(cursor_class, connection, **kwargs): + return cursor_class( + connection=connection, + converter=MagicMock(), + formatter=MagicMock(), + retry_config=RetryConfig(), + s3_staging_dir=None, + schema_name=None, + catalog_name=None, + work_group="spark", + poll_interval=0, + encryption_option=None, + kms_key=None, + kill_on_interrupt=False, + result_reuse_enable=False, + result_reuse_minutes=60, + **kwargs, + ) + + class TestSparkBaseCursor: @pytest.mark.parametrize( "state", @@ -97,3 +134,105 @@ def test_exists_session_raises_on_failure_state(self): ): cursor._exists_session("session_id") cursor._connection.client.get_session.assert_called_once_with(SessionId="session_id") + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + def test_init_starts_session(self, cursor_class): + connection = _connection() + with patch.object(SparkBaseCursor, "_wait_for_idle_session"): + cursor = _init_cursor(cursor_class, connection) + + assert cursor.session_id == "new-session" + connection.client.terminate_session.assert_not_called() + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + @pytest.mark.parametrize( + "error", + [OperationalError("Session did not become idle."), KeyboardInterrupt()], + ) + def test_init_terminates_new_session_that_does_not_become_idle(self, cursor_class, error): + connection = _connection() + with ( + patch.object(SparkBaseCursor, "_wait_for_idle_session", side_effect=error), + pytest.raises(type(error)) as exc_info, + ): + _init_cursor(cursor_class, connection) + + assert exc_info.value is error + connection.client.terminate_session.assert_called_once_with(SessionId="new-session") + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + @pytest.mark.parametrize( + "state", + [ + AthenaSessionStatus.STATE_TERMINATED, + AthenaSessionStatus.STATE_DEGRADED, + AthenaSessionStatus.STATE_FAILED, + ], + ) + def test_init_terminates_new_session_in_failure_state(self, cursor_class, state): + connection = _connection() + connection.client.get_session_status.return_value = { + "SessionId": "new-session", + "Status": {"State": state, "StateChangeReason": "session failure reason"}, + } + with ( + patch("pyathena.spark.common.time.sleep", side_effect=AssertionError("slept")), + pytest.raises( + OperationalError, + match=rf"^Session: new-session is {state}\. session failure reason$", + ), + ): + _init_cursor(cursor_class, connection) + + connection.client.get_session_status.assert_called_once_with(SessionId="new-session") + connection.client.terminate_session.assert_called_once_with(SessionId="new-session") + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + def test_init_keeps_original_error_when_cleanup_fails(self, cursor_class, caplog): + connection = _connection() + connection.client.terminate_session.side_effect = ClientError( + {"Error": {"Code": "InternalServerException", "Message": "Cleanup failed."}}, + "TerminateSession", + ) + error = OperationalError("Session did not become idle.") + with ( + caplog.at_level(logging.ERROR, logger="pyathena.spark.common"), + patch.object(SparkBaseCursor, "_wait_for_idle_session", side_effect=error), + pytest.raises(OperationalError) as exc_info, + ): + _init_cursor(cursor_class, connection) + + assert exc_info.value is error + connection.client.terminate_session.assert_called_once_with(SessionId="new-session") + assert "Failed to terminate session: new-session." in caplog.text + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + def test_init_does_not_terminate_supplied_session(self, cursor_class): + connection = _connection() + error = OperationalError("Session did not become idle.") + with ( + patch.object(SparkBaseCursor, "_wait_for_idle_session", side_effect=error), + pytest.raises(OperationalError) as exc_info, + ): + _init_cursor(cursor_class, connection, session_id="supplied-session") + + assert exc_info.value is error + connection.client.start_session.assert_not_called() + connection.client.terminate_session.assert_not_called() + + @pytest.mark.parametrize("cursor_class", SPARK_CURSOR_CLASSES) + def test_init_does_not_start_session_when_s3_client_fails(self, cursor_class): + connection = _connection() + connection.session.client.side_effect = ValueError("Invalid S3 client configuration.") + with pytest.raises(ValueError, match=r"^Invalid S3 client configuration\.$"): + _init_cursor(cursor_class, connection) + + connection.client.start_session.assert_not_called() + connection.client.terminate_session.assert_not_called() + + def test_async_init_does_not_start_session_when_executor_fails(self): + connection = _connection() + with pytest.raises(ValueError, match="max_workers must be greater than 0"): + _init_cursor(AsyncSparkCursor, connection, max_workers=0) + + connection.client.start_session.assert_not_called() From d2240e3a04f437132d554c94c30e8e556122000c Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 20:47:04 +0900 Subject: [PATCH 6/7] Backport #826: Cast Double and DOUBLE_PRECISION to DOUBLE SQLAlchemy's generic Double and DOUBLE_PRECISION are not DOUBLE subclasses, so CAST fell through to the Float branch and rendered REAL, losing precision. This backports only the compiler fix from #826 (cherry picked from commit 3ddb868f85c0a168253ca15fc5286f8970c7392a), with a regression test that runs on SQLAlchemy 2.0. Co-Authored-By: Claude Opus 5.5 --- pyathena/sqlalchemy/compiler.py | 2 +- tests/pyathena/sqlalchemy/test_compiler.py | 21 +++++++++++++++++++++ 2 files changed, 22 insertions(+), 1 deletion(-) diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 6a119eb34..72223dbc5 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -296,7 +296,7 @@ def visit_cast(self, cast: Cast[Any], **kwargs): type_clause = "CHAR" elif isinstance(cast.type, (types.BINARY, types.VARBINARY)): type_clause = "VARBINARY" - elif hasattr(types, "DOUBLE") and isinstance(cast.type, types.DOUBLE): + elif hasattr(types, "Double") and isinstance(cast.type, types.Double): type_clause = "DOUBLE" elif isinstance(cast.type, (types.FLOAT, types.Float, types.REAL)): # https://docs.aws.amazon.com/athena/latest/ug/data-types.html diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index 1fb120f5a..cfe310106 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -10,9 +10,12 @@ Numeric, String, Table, + cast, + column, exc, func, select, + types, ) from sqlalchemy.engine.url import make_url from sqlalchemy.sql import literal, literal_column @@ -137,6 +140,24 @@ def test_visit_json(self): result = compiler.visit_JSON(json_type) assert result == "JSON" + @pytest.mark.skipif(not hasattr(types, "Double"), reason="Requires SQLAlchemy 2.0") + @pytest.mark.parametrize( + ("type_name", "ddl", "cast_type"), + [ + ("Float", "FLOAT", "REAL"), + ("FLOAT", "FLOAT", "REAL"), + ("REAL", "FLOAT", "REAL"), + ("Double", "DOUBLE", "DOUBLE"), + ("DOUBLE", "DOUBLE", "DOUBLE"), + ("DOUBLE_PRECISION", "DOUBLE", "DOUBLE"), + ], + ) + def test_floating_point_types(self, type_name, ddl, cast_type): + dialect = AthenaDialect() + type_ = getattr(types, type_name)() + assert dialect.type_compiler_instance.process(type_) == ddl + assert str(cast(column("x"), type_).compile(dialect=dialect)) == f"CAST(x AS {cast_type})" + class TestAthenaStatementCompiler: """Test cases for Athena statement compiler functionality.""" From 610d6c3ae7b954de46ca5b5cea86c89050cbf67f Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 20:48:45 +0900 Subject: [PATCH 7/7] Backport #870: Render Hive STRUCT syntax in table column DDL CREATE TABLE rendered AthenaStruct columns as ROW(...), which Athena rejects in table DDL. That affected top-level STRUCT columns and STRUCT values inside MAP and ARRAY. Column DDL now renders STRUCT at every nesting depth, and quotes field names with the DDL identifier preparer. This reimplements #870 (cherry picked from commit cf673c2f9118b98797ddcb5e50beba7f8a70dd64) for 3.x, which lacks the ARRAY DDL context from #774. Integer spelling inside MAP and STRUCT stays INTEGER, which Athena's DDL accepts. Direct type compilation, CAST output, and an empty AthenaStruct() are unchanged. Co-Authored-By: Claude Opus 5.5 --- docs/sqlalchemy.md | 8 ++- pyathena/sqlalchemy/compiler.py | 42 +++++++++--- tests/pyathena/sqlalchemy/test_base.py | 63 ++++++++++++++--- tests/pyathena/sqlalchemy/test_compiler.py | 80 +++++++++++++++++++++- 4 files changed, 173 insertions(+), 20 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index b7be6fe7c..6740caa35 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -697,12 +697,14 @@ This generates the following SQL structure: ```sql CREATE TABLE users ( - id INTEGER, - profile ROW(name STRING, age INTEGER, email STRING), - settings ROW(theme STRING, notifications ROW(email STRING, push STRING)) + id INT, + profile STRUCT, + settings STRUCT> ) ``` +`CREATE TABLE` renders `AthenaStruct` columns with Hive `STRUCT` syntax at every nesting depth, including STRUCT values inside MAP and ARRAY. + #### Querying STRUCT data PyAthena automatically converts STRUCT data between different formats: diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 72223dbc5..e68143311 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -171,22 +171,47 @@ def visit_tinyint(self, type_, **kw): def visit_enum(self, type_, **kw): return self.visit_string(type_, **kw) + def _enable_hive_column_ddl(self, kw: dict[str, Any]) -> bool: + """Enable Hive type syntax for a CREATE TABLE column type. + + ``get_column_specification`` passes the column as ``type_expression``. + The flag set here is passed to nested types, so STRUCT at any depth + of a column type uses ``STRUCT``. Direct type compilation + and CAST leave it unset. + + Args: + kw: Type-compiler keyword arguments. When Hive syntax applies, + ``_athena_hive_ddl`` is set so nested types keep it. + + Returns: + True when the type should use Hive DDL syntax. + """ + if kw.get("_athena_hive_ddl") or isinstance(kw.get("type_expression"), Column): + kw["_athena_hive_ddl"] = True + return True + return False + def visit_struct(self, type_, **kw): - if isinstance(type_, AthenaStruct): - if type_.fields: - field_specs = [] - for field_name, field_type in type_.fields.items(): - field_type_str = self.process(field_type, **kw) - field_specs.append(f"{field_name} {field_type_str}") - return f"ROW({', '.join(field_specs)})" + if not isinstance(type_, AthenaStruct) or not type_.fields: return "ROW()" - return "ROW()" + if self._enable_hive_column_ddl(kw): + preparer = AthenaDDLIdentifierPreparer(self.dialect) + fields = ", ".join( + f"{preparer.quote(name)}:{self.process(field_type, **kw)}" + for name, field_type in type_.fields.items() + ) + return f"STRUCT<{fields}>" + fields = ", ".join( + f"{name} {self.process(field_type, **kw)}" for name, field_type in type_.fields.items() + ) + return f"ROW({fields})" def visit_STRUCT(self, type_, **kw): return self.visit_struct(type_, **kw) def visit_map(self, type_, **kw): if isinstance(type_, AthenaMap): + self._enable_hive_column_ddl(kw) key_type_str = self.process(type_.key_type, **kw) value_type_str = self.process(type_.value_type, **kw) return f"MAP<{key_type_str}, {value_type_str}>" @@ -197,6 +222,7 @@ def visit_MAP(self, type_, **kw): def visit_array(self, type_, **kw): if isinstance(type_, AthenaArray): + self._enable_hive_column_ddl(kw) item_type_str = self.process(type_.item_type, **kw) return f"ARRAY<{item_type_str}>" return "ARRAY" diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 7c0952e94..5917950aa 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -2389,7 +2389,7 @@ def test_create_table_with_array_types(self, engine): assert "tags ARRAY" in ddl_string assert "scores ARRAY" in ddl_string assert "nested_arrays ARRAY>" in ddl_string - assert "struct_array ARRAY" in ddl_string + assert "struct_array ARRAY>" in ddl_string def test_create_table_with_map_types(self, engine): """Test DDL compilation for MAP types.""" @@ -2419,7 +2419,7 @@ def test_create_table_with_map_types(self, engine): # Verify MAP types are correctly compiled assert "attributes MAP" in ddl_string assert "metrics MAP" in ddl_string - assert "complex_map MAP" in ddl_string + assert "complex_map MAP>" in ddl_string assert "nested_map MAP>" in ddl_string def test_create_table_with_struct_types(self, engine): @@ -2461,12 +2461,12 @@ def test_create_table_with_struct_types(self, engine): ddl_string = str(create_ddl) # Verify STRUCT types are correctly compiled - assert "user_info ROW(name STRING, age INTEGER, email STRING)" in ddl_string + assert "user_info STRUCT" in ddl_string assert ( - "nested_struct ROW(personal ROW(first_name STRING, last_name STRING), " - "preferences MAP)" in ddl_string + "nested_struct STRUCT, " + "preferences:MAP>" in ddl_string ) - assert "struct_with_array ROW(tags ARRAY, scores ARRAY)" in ddl_string + assert "struct_with_array STRUCT, scores:ARRAY>" in ddl_string def test_create_table_with_complex_nested_types(self, engine): """Test DDL compilation for complex nested combinations of ARRAY, MAP, and STRUCT.""" @@ -2499,11 +2499,58 @@ def test_create_table_with_complex_nested_types(self, engine): # Verify complex nested type is correctly compiled expected_type = ( - "data ARRAY, " - "tags ARRAY)>>" + "data ARRAY, " + "tags:ARRAY>>>" ) assert expected_type in ddl_string + def test_external_parquet_struct_columns_round_trip(self, engine): + """Create a Parquet table of top-level and MAP-nested STRUCTs and read the fields back.""" + _, conn = engine + table_name = "test_external_parquet_struct_columns" + table = Table( + table_name, + MetaData(schema=ENV.schema), + Column( + "profile", + AthenaStruct( + ("name", types.String), + ("age", types.Integer), + ( + "address", + AthenaStruct(("city", types.String), ("zip", types.Integer)), + ), + ), + ), + Column( + "labels", + AthenaMap( + types.String, + AthenaStruct(("value", types.String), ("count", types.Integer)), + ), + ), + awsathena_location=f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/", + awsathena_file_format="PARQUET", + ) + table.create(bind=conn) + conn.execute( + text( + f"INSERT INTO {ENV.schema}.{table_name} VALUES (" + "CAST(ROW('Ada', 36, ROW('London', 12345)) AS " + "ROW(name VARCHAR, age INTEGER, address ROW(city VARCHAR, zip INTEGER))), " + "MAP(ARRAY['home'], ARRAY[CAST(ROW('Lovelace', 2) AS " + "ROW(value VARCHAR, count INTEGER))]))" + ) + ) + row = conn.execute( + text( + "SELECT profile.name, profile.age, profile.address.city, " + "profile.address.zip, labels['home'].value, labels['home'].count " + f"FROM {ENV.schema}.{table_name}" + ) + ).one() + assert tuple(row) == ("Ada", 36, "London", 12345, "Lovelace", 2) + def test_sqlalchemy_execute_with_execution_options_callback(self, engine): """Test callback functionality through SQLAlchemy execution_options.""" engine, conn = engine diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index cfe310106..37b46d099 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -297,7 +297,10 @@ def test_visit_truediv_binary(self, expression, expected): class TestAthenaDDLCompiler: - """Compile-only (no AWS) tests for the DDL compiler's S3 Tables support. + """Compile-only (no AWS) tests for the DDL compiler. + + Covers column type rendering in CREATE TABLE, where STRUCT at any depth + uses Hive syntax, and S3 Tables support. S3 Tables are queried by setting the connection ``catalog_name`` to ``s3tablescatalog/`` and using the namespace as the table @@ -315,6 +318,81 @@ def _s3tables_dialect(self, **connect_opts): } return dialect + def _ddl(self, *columns): + table = Table( + "events", + MetaData(schema="analytics"), + *columns, + awsathena_location="s3://bucket/events/", + awsathena_file_format="PARQUET", + ) + return str(CreateTable(table).compile(dialect=AthenaDialect())) + + def test_create_table_renders_hive_struct_syntax(self): + ddl = self._ddl( + Column("id", Integer), + Column( + "profile", + AthenaStruct( + ("name", String), + ("age", Integer), + ("address", AthenaStruct(("city", String), ("zip", Integer))), + ), + ), + Column( + "labels", + AthenaMap(String, AthenaStruct(("value", String), ("count", Integer))), + ), + Column( + "nested_maps", + AthenaMap(String, AthenaMap(Integer, AthenaStruct(("n", Integer)))), + ), + Column( + "mixed", + AthenaStruct( + ("tags", AthenaArray(String)), + ("attrs", AthenaMap(String, Integer)), + ), + ), + Column("items", AthenaArray(AthenaStruct(("name", String), ("qty", Integer)))), + Column( + "deep", + AthenaArray(AthenaMap(String, AthenaStruct(("flag", types.Boolean)))), + ), + ) + assert "id INT" in ddl + assert ( + "profile STRUCT>" + ) in ddl + assert "labels MAP>" in ddl + assert "nested_maps MAP>>" in ddl + assert "mixed STRUCT, attrs:MAP>" in ddl + assert "items ARRAY>" in ddl + assert "deep ARRAY>>" in ddl + assert "ROW(" not in ddl + + def test_struct_field_quoting_follows_ddl_preparer(self): + ddl = self._ddl( + Column( + "payload", + AthenaStruct(("date", String), ("a`b", String), ("first name", String)), + ) + ) + assert "payload STRUCT<`date`:STRING, `a``b`:STRING, `first name`:STRING>" in ddl + + def test_struct_type_without_column_context_stays_row(self): + struct_type = AthenaStruct(("name", String), ("tags", AthenaArray(String))) + map_type = AthenaMap(String, AthenaStruct(("n", Integer))) + array_type = AthenaArray(AthenaStruct(("n", Integer))) + compiler = AthenaDialect().type_compiler_instance + assert compiler.process(struct_type) == "ROW(name STRING, tags ARRAY)" + assert compiler.process(map_type) == "MAP" + assert compiler.process(array_type) == "ARRAY" + + def test_empty_struct_column_stays_row(self): + ddl = self._ddl(Column("empty", AthenaStruct())) + assert "empty ROW()" in ddl + def test_create_table_s3tables_catalog_omits_location(self): table = Table( "tbl",