diff --git a/.github/workflows/release.yaml b/.github/workflows/release.yaml index d7918824..0ca9581a 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 535bf9a8..22be768a 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 d8e6a3ee..1ab76af2 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 }} diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index b7be6fe7..6740caa3 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/pandas/util.py b/pyathena/pandas/util.py index b2b96612..9337ce32 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/pyathena/spark/async_cursor.py b/pyathena/spark/async_cursor.py index 1a476f20..17aa2a3e 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,12 +104,24 @@ 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: - 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/pyathena/spark/common.py b/pyathena/spark/common.py index d25c0b0f..7ffbb7bb 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 @@ -126,21 +144,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, @@ -161,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, @@ -181,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, @@ -195,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/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 6a119eb3..e6814331 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" @@ -296,7 +322,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/pandas/test_util.py b/tests/pyathena/pandas/test_util.py index 788308da..1acbc07e 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, diff --git a/tests/pyathena/spark/test_async_cursor.py b/tests/pyathena/spark/test_async_cursor.py index e921bdc3..6f9f2d25 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 diff --git a/tests/pyathena/spark/test_common.py b/tests/pyathena/spark/test_common.py new file mode 100644 index 00000000..9dbfd27f --- /dev/null +++ b/tests/pyathena/spark/test_common.py @@ -0,0 +1,238 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# 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( + {"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 + + +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", + [ + 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") + + @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() diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 7c0952e9..5917950a 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 1fb120f5..37b46d09 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.""" @@ -276,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 @@ -294,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",