From 0fac8b1cf27af8545bb67a927b75d36d9681d6e8 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:32:20 +0900 Subject: [PATCH 01/20] Backport #769: Enable Iceberg UPDATE target subquery compliance (cherry picked from commit 714efbba231c6401f940ca4ccf8e6480a24f723d) Co-Authored-By: Claude Opus 5.5 --- pyathena/sqlalchemy/requirements.py | 3 ++- tests/sqlalchemy/test_suite.py | 30 +++++++++++++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) diff --git a/pyathena/sqlalchemy/requirements.py b/pyathena/sqlalchemy/requirements.py index e600a949..8e5e2f34 100644 --- a/pyathena/sqlalchemy/requirements.py +++ b/pyathena/sqlalchemy/requirements.py @@ -72,7 +72,8 @@ def duplicate_key_raises_integrity_error(self): @property def update_where_target_in_subquery(self): - return unsupported() + # Verified with Iceberg tables on Athena engine version 3. + return supported() @property def recursive_fk_cascade(self): diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 10be6405..763bd2e3 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -1,9 +1,12 @@ import pytest +from sqlalchemy import func, select, testing +from sqlalchemy.testing import eq_ from sqlalchemy.testing.suite import * # noqa: F403 from sqlalchemy.testing.suite import FetchLimitOffsetTest as _FetchLimitOffsetTest from sqlalchemy.testing.suite import HasTableTest as _HasTableTest from sqlalchemy.testing.suite import InsertBehaviorTest as _InsertBehaviorTest from sqlalchemy.testing.suite import IntegerTest as _IntegerTest +from sqlalchemy.testing.suite import SimpleUpdateDeleteTest as _SimpleUpdateDeleteTest from sqlalchemy.testing.suite import StringTest as _StringTest del BinaryTest # noqa: F821 @@ -26,6 +29,33 @@ del UuidTest # noqa: F821 +class SimpleUpdateDeleteTest(_SimpleUpdateDeleteTest): + @testing.variation("criteria", ["rows", "norows", "aggregate"]) + @testing.requires.update_where_target_in_subquery + def test_update_where_target_in_subquery(self, connection, criteria): + t = self.tables.plain_pk + if criteria.rows: + subquery = select(t.c.id).where(t.c.id < 3) + expected = [(1, "updated"), (2, "updated"), (3, "d3")] + rowcount = 2 + elif criteria.norows: + subquery = select(t.c.id).where(t.c.id < 0) + expected = [(1, "d1"), (2, "d2"), (3, "d3")] + rowcount = 0 + elif criteria.aggregate: + subquery = select(func.max(t.c.id)) + expected = [(1, "d1"), (2, "d2"), (3, "updated")] + rowcount = 1 + else: + criteria.fail() + + r = connection.execute(t.update().where(t.c.id.in_(subquery)), {"data": "updated"}) + assert not r.is_insert + assert not r.returns_rows + assert r.rowcount == rowcount + eq_(connection.execute(t.select().order_by(t.c.id)).fetchall(), expected) + + class HasTableTest(_HasTableTest): @pytest.mark.skip("No cache is used when creating tables.") def test_has_table_cache(self, metadata): From 759299344b1a2675cf1b7958403ca1d149acb433 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:32:21 +0900 Subject: [PATCH 02/20] Backport #772: Restore SQLAlchemy CTE compliance coverage (cherry picked from commit f7b23698bb711c7d02b892d5f0295f02f1885e04) Conflict resolution: tests/sqlalchemy/test_suite.py conflicted because #771 and #777 (not backported) changed its imports and helpers. Applied only #772's own changes: import CTETest, drop its del, and add the CTETest subclass. Also added Integer to the existing sqlalchemy import, which the subclass uses and which master had imported through #771. Co-Authored-By: Claude Opus 5.5 --- pyathena/sqlalchemy/requirements.py | 27 +++++++++++++++++++++++++++ tests/sqlalchemy/test_suite.py | 12 ++++++++++-- 2 files changed, 37 insertions(+), 2 deletions(-) diff --git a/pyathena/sqlalchemy/requirements.py b/pyathena/sqlalchemy/requirements.py index 8e5e2f34..a9ccc0f4 100644 --- a/pyathena/sqlalchemy/requirements.py +++ b/pyathena/sqlalchemy/requirements.py @@ -103,8 +103,35 @@ def window_functions(self): @property def ctes(self): + # Recursive CTEs require Athena engine version 3 and have a maximum depth of 10. return supported() + @property + def ctes_with_values(self): + return supported() + + @property + def ctes_with_update_delete(self): + return exclusions.skip_if( + lambda _: True, "Athena does not support WITH preceding UPDATE or DELETE." + ) + + @property + def ctes_on_dml(self): + return exclusions.skip_if( + lambda _: True, "Athena does not support INSERT, UPDATE, or DELETE inside a CTE." + ) + + @property + def update_from(self): + return exclusions.skip_if(lambda _: True, "Athena does not support UPDATE ... FROM.") + + @property + def delete_from(self): + return exclusions.skip_if( + lambda _: True, "Athena does not support DELETE ... USING or multi-table DELETE." + ) + @property def views(self): return supported() diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 763bd2e3..dc67d575 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -1,7 +1,8 @@ import pytest -from sqlalchemy import func, select, testing +from sqlalchemy import Integer, func, select, testing from sqlalchemy.testing import eq_ from sqlalchemy.testing.suite import * # noqa: F403 +from sqlalchemy.testing.suite import CTETest as _CTETest from sqlalchemy.testing.suite import FetchLimitOffsetTest as _FetchLimitOffsetTest from sqlalchemy.testing.suite import HasTableTest as _HasTableTest from sqlalchemy.testing.suite import InsertBehaviorTest as _InsertBehaviorTest @@ -13,7 +14,6 @@ del ComponentReflectionTest # noqa: F821 del ComponentReflectionTestExtra # noqa: F821 del CompositeKeyReflectionTest # noqa: F821 -del CTETest # noqa: F821 del DateTimeMicrosecondsTest # noqa: F821 del DifficultParametersTest # noqa: F821 del DistinctOnTest # noqa: F821 @@ -29,6 +29,14 @@ del UuidTest # noqa: F821 +class CTETest(_CTETest): + @classmethod + def define_tables(cls, metadata): + super().define_tables(metadata) + # The suite removes unsupported foreign keys, so parent_id cannot infer its type. + metadata.tables["some_table"].c.parent_id.type = Integer() + + class SimpleUpdateDeleteTest(_SimpleUpdateDeleteTest): @testing.variation("criteria", ["rows", "norows", "aggregate"]) @testing.requires.update_where_target_in_subquery From 311ce59b2ad24776457af1e7d563625ca42bd1df Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:34:57 +0900 Subject: [PATCH 03/20] Backport #770: Restore SQLAlchemy binary compliance and preserve CSV binary NULLs (cherry picked from commit 2cd0bdfdafd66d75b496f283af0dbfc992a5faf0) Conflict resolution: tests/sqlalchemy/test_suite.py conflicted because #771 and #777 (not backported) changed its imports and helpers. Applied only #770's own changes: import BinaryTest, MetaData, types, and Table as SATable; drop the BinaryTest del; and add the BinaryTest subclass verbatim. The subclass uses the 'sa_testing' alias for sqlalchemy.testing, which master introduced in #771, so the import is added alongside 3.x's existing 'testing' import. Later backports then apply master's code unchanged. Co-Authored-By: Claude Opus 5.5 --- docs/null_handling.md | 11 + pyathena/__init__.py | 1 + pyathena/aio/sqlalchemy/base.py | 2 + pyathena/arrow/result_set.py | 13 +- pyathena/formatter.py | 10 + pyathena/pandas/cursor.py | 11 +- pyathena/pandas/reader.py | 64 ++++++ pyathena/pandas/result_set.py | 240 +++++++++++++++++---- pyathena/s3fs/reader.py | 21 +- pyathena/sqlalchemy/base.py | 4 + pyathena/sqlalchemy/compiler.py | 2 +- pyathena/sqlalchemy/types.py | 10 + tests/pyathena/aio/arrow/test_cursor.py | 16 ++ tests/pyathena/aio/pandas/test_cursor.py | 15 ++ tests/pyathena/aio/sqlalchemy/test_base.py | 60 +++++- tests/pyathena/aio/test_cursor.py | 20 +- tests/pyathena/arrow/test_async_cursor.py | 16 ++ tests/pyathena/arrow/test_cursor.py | 20 ++ tests/pyathena/conftest.py | 3 + tests/pyathena/pandas/test_async_cursor.py | 15 ++ tests/pyathena/pandas/test_cursor.py | 191 ++++++++++++++++ tests/pyathena/pandas/test_reader.py | 38 ++++ tests/pyathena/sqlalchemy/test_base.py | 61 ++++++ tests/pyathena/test_cursor.py | 31 ++- tests/pyathena/test_formatter.py | 18 ++ tests/sqlalchemy/test_suite.py | 35 ++- 26 files changed, 867 insertions(+), 61 deletions(-) create mode 100644 pyathena/pandas/reader.py create mode 100644 tests/pyathena/pandas/test_reader.py diff --git a/docs/null_handling.md b/docs/null_handling.md index 3032c162..3b28700d 100644 --- a/docs/null_handling.md +++ b/docs/null_handling.md @@ -60,6 +60,17 @@ which correctly interprets unquoted empty values as NULL, while `S3FSCursor` use `AthenaCSVReader` that respects CSV quoting rules. ``` +## Binary Values + +The string comparison above does not apply to `VARBINARY` columns. +With the default CSV settings and converters, pandas and Arrow cursors distinguish SQL NULL from empty binary +values when reading CSV results: `fetchone()`, `fetchmany()`, and `fetchall()` return `None` +for NULL and `b''` for an empty binary value. +This also applies to their asynchronous variants and pandas chunked reads. +Pandas DataFrames preserve the same values. +Arrow Tables retain the CSV hexadecimal strings, with NULL represented as an Arrow null; +fetch methods convert the hexadecimal strings to Python bytes. + ## Default Cursor (API-based) The default `Cursor` and `DictCursor` fetch results directly from the Athena API, diff --git a/pyathena/__init__.py b/pyathena/__init__.py index e462b28b..acba8de5 100644 --- a/pyathena/__init__.py +++ b/pyathena/__init__.py @@ -63,6 +63,7 @@ def __hash__(self): Date: type[datetime.date] = datetime.date Time: type[datetime.time] = datetime.time Timestamp: type[datetime.datetime] = datetime.datetime +Binary: type[bytes] = bytes @overload diff --git a/pyathena/aio/sqlalchemy/base.py b/pyathena/aio/sqlalchemy/base.py index a0308937..9f692e40 100644 --- a/pyathena/aio/sqlalchemy/base.py +++ b/pyathena/aio/sqlalchemy/base.py @@ -163,6 +163,8 @@ class AsyncAdapt_pyathena_dbapi: """ paramstyle = "pyformat" + Binary = pyathena.Binary + BINARY = pyathena.BINARY # DBAPI exception hierarchy Error = Error diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index c190a98b..2ff47683 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -276,6 +276,7 @@ def _read_csv(self) -> Table: ): return pa.Table.from_pydict({}) length = self._get_content_length() + binary_columns = {d[0] for d in self.description or [] if d[1] == "varbinary"} if length and self.output_location.endswith(".txt"): description = self.description if self.description else [] column_names = [d[0] for d in description] @@ -296,6 +297,7 @@ def _read_csv(self) -> Table: parse_opts = csv.ParseOptions( delimiter=",", quote_char='"', + ignore_empty_lines=not binary_columns, double_quote=True, escape_char=False, ) @@ -304,16 +306,25 @@ def _read_csv(self) -> Table: bucket, key = parse_output_location(self.output_location) try: - return csv.read_csv( + table = csv.read_csv( self._fs.open_input_stream(f"{bucket}/{key}"), read_options=read_opts, parse_options=parse_opts, convert_options=csv.ConvertOptions( + strings_can_be_null=bool(binary_columns), quoted_strings_can_be_null=False, timestamp_parsers=self.timestamp_parsers, column_types=self.column_types, ), ) + if binary_columns: + for index, field in enumerate(table.schema): + if field.name not in binary_columns and ( + pa.types.is_string(field.type) or pa.types.is_binary(field.type) + ): + # Preserve the existing CSV behavior for non-binary Athena columns. + table = table.set_column(index, field, table.column(index).fill_null("")) + return table except Exception as e: _logger.exception(f"Failed to read {bucket}/{key}.") raise OperationalError(*e.args) from e diff --git a/pyathena/formatter.py b/pyathena/formatter.py index 89564ea4..bb305f7f 100644 --- a/pyathena/formatter.py +++ b/pyathena/formatter.py @@ -251,6 +251,12 @@ def _format_str(formatter: Formatter, escaper: Callable[[str], str], val: Any) - return escaper(val) +def _format_binary( + formatter: Formatter, escaper: Callable[[str], str], val: bytes | bytearray | memoryview +) -> str: + return f"X'{val.hex()}'" + + def _format_seq(formatter: Formatter, escaper: Callable[[str], str], val: Any) -> Any: results = [] for v in val: @@ -291,6 +297,9 @@ def _format_decimal(formatter: Formatter, escaper: Callable[[str], str], val: An Decimal: _format_decimal, bool: _format_bool, str: _format_str, + bytes: _format_binary, + bytearray: _format_binary, + memoryview: _format_binary, list: _format_seq, set: _format_seq, tuple: _format_seq, @@ -307,6 +316,7 @@ class DefaultParameterFormatter(Formatter): Supported types: - None: Converts to SQL NULL - Strings: Properly escaped and quoted + - Binary data: bytes, bytearray, memoryview as hexadecimal literals - Numbers: int, float, Decimal - Dates and times: date, datetime, time - Booleans: Converted to SQL boolean literals diff --git a/pyathena/pandas/cursor.py b/pyathena/pandas/cursor.py index edb0f356..883172c0 100644 --- a/pyathena/pandas/cursor.py +++ b/pyathena/pandas/cursor.py @@ -294,9 +294,10 @@ def iter_chunks(self) -> Generator[DataFrame, None, None]: import gc - for chunk_count, chunk in enumerate(result_set.iter_chunks(), 1): - yield chunk + with result_set.iter_chunks() as chunks: + for chunk_count, chunk in enumerate(chunks, 1): + yield chunk - # Suggest garbage collection every 10 chunks for large datasets - if chunk_count % 10 == 0: - gc.collect() + # Suggest garbage collection every 10 chunks for large datasets + if chunk_count % 10 == 0: + gc.collect() diff --git a/pyathena/pandas/reader.py b/pyathena/pandas/reader.py new file mode 100644 index 00000000..e8e1a2c4 --- /dev/null +++ b/pyathena/pandas/reader.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +import re +from io import RawIOBase +from typing import Any + +from pyathena.s3fs.reader import AthenaCSVReader + +_BINARY_NULL = "__PYATHENA_BINARY_NULL__" +_CSV_FIELD = re.compile(r'(?:^|,)(?P"[^"]*(?:""[^"]*)*"|[^,]*)') + + +class BinaryCSVReader(RawIOBase): + """Preserve binary NULL fields before pandas discards CSV quoting information. + + The marker cannot occur in Athena's hexadecimal encoding of binary values. + Only unquoted empty binary fields are rewritten; all other CSV text is + passed through unchanged. Records are streamed to support chunked reads. + """ + + def __init__(self, stream: Any, binary_columns: set[int]) -> None: + super().__init__() + self._reader = AthenaCSVReader(stream) + self._binary_columns = binary_columns + self._header = True + self._buffer = b"" + + def readable(self) -> bool: + return True + + def readinto(self, buffer: Any) -> int: + if self.closed: + raise ValueError("I/O operation on closed file.") + if not len(buffer): + return 0 + if not self._buffer: + try: + record = self._reader._read_record() + except StopIteration: + self._reader.close() + return 0 + if self._header: + self._header = False + else: + parts: list[str] = [] + start = 0 + for index, field in enumerate(_CSV_FIELD.finditer(record.rstrip("\r\n"))): + if index in self._binary_columns and not field.group("value"): + pos = field.start("value") + parts.extend((record[start:pos], _BINARY_NULL)) + start = pos + parts.append(record[start:]) + record = "".join(parts) + self._buffer = record.encode("utf-8") + size = min(len(buffer), len(self._buffer)) + buffer[:size] = self._buffer[:size] + self._buffer = self._buffer[size:] + return size + + def close(self) -> None: + try: + self._reader.close() + finally: + super().close() diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 7b4e1136..4eba6d60 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -1,8 +1,12 @@ from __future__ import annotations +import csv import logging from collections import abc from collections.abc import Callable, Iterable, Iterator +from contextlib import ExitStack +from functools import partial +from io import BufferedReader, StringIO, TextIOWrapper from multiprocessing import cpu_count from typing import ( TYPE_CHECKING, @@ -10,10 +14,13 @@ ClassVar, ) +from fsspec import open as filesystem_open + from pyathena import OperationalError from pyathena.converter import Converter from pyathena.error import ProgrammingError from pyathena.model import AthenaQueryExecution +from pyathena.pandas.reader import _BINARY_NULL, BinaryCSVReader from pyathena.result_set import AthenaResultSet from pyathena.util import RetryConfig, parse_output_location @@ -26,6 +33,10 @@ _logger = logging.getLogger(__name__) +def _convert_binary_csv(converter: Callable[[str | None], Any], value: str) -> Any: + return converter(None if value == _BINARY_NULL else value) + + def _no_trunc_date(df: DataFrame) -> DataFrame: return df @@ -59,12 +70,14 @@ def __init__( self, reader: TextFileReader | DataFrame, trunc_date: Callable[[DataFrame], DataFrame], + csv_stream: TextIOWrapper | None = None, ) -> None: """Initialize the iterator. Args: reader: Either a TextFileReader (for chunked) or a single DataFrame. trunc_date: Function to apply date truncation to each chunk. + csv_stream: Optional CSV stream owned and closed by this iterator. """ from pandas import DataFrame @@ -73,6 +86,7 @@ def __init__( else: self._reader = reader self._trunc_date = trunc_date + self._csv_stream = csv_stream def __next__(self) -> DataFrame: """Get the next DataFrame chunk. @@ -86,7 +100,7 @@ def __next__(self) -> DataFrame: try: df = next(self._reader) return self._trunc_date(df) - except StopIteration: + except BaseException: self.close() raise @@ -106,8 +120,14 @@ def close(self) -> None: """Close the iterator and release resources.""" from pandas.io.parsers import TextFileReader - if isinstance(self._reader, TextFileReader): - self._reader.close() + reader = self._reader + self._reader = iter(()) + try: + if isinstance(reader, TextFileReader): + reader.close() + finally: + if self._csv_stream is not None: + self._csv_stream.close() def iterrows(self) -> Iterator[tuple[int, dict[str, Any]]]: """Iterate over rows as (index, row_dict) tuples. @@ -137,9 +157,13 @@ def get_chunk(self, size: int | None = None) -> DataFrame: """ from pandas.io.parsers import TextFileReader - if isinstance(self._reader, TextFileReader): - return self._reader.get_chunk(size) - return next(self._reader) + try: + if isinstance(self._reader, TextFileReader): + return self._reader.get_chunk(size) + return next(self._reader) + except BaseException: + self.close() + raise def as_pandas(self) -> DataFrame: """Collect all chunks into a single DataFrame. @@ -281,6 +305,7 @@ def __init__( self._data_manifest: list[str] = [] self._kwargs = kwargs self._fs = self.__s3_file_system() + self._csv_stream: TextIOWrapper | None = None # Cache time column names for efficient _trunc_date processing description = self.description if self.description else [] @@ -291,7 +316,7 @@ def __init__( if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: df = self._as_pandas() trunc_date = _no_trunc_date if self.is_unload else self._trunc_date - self._df_iter = PandasDataFrameIterator(df, trunc_date) + self._df_iter = PandasDataFrameIterator(df, trunc_date, self._csv_stream) elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: df = self._as_pandas_from_api() self._df_iter = PandasDataFrameIterator(df, self._trunc_date) @@ -504,18 +529,6 @@ def _read_csv(self) -> TextFileReader | DataFrame: if length == 0: return pd.DataFrame() - if self.output_location.endswith(".txt"): - sep = "\t" - header = None - description = self.description if self.description else [] - names = [d[0] for d in description] - elif self.output_location.endswith(".csv"): - sep = "," - header = 0 - names = None - else: - return pd.DataFrame() - # Chunksize determination with user preference priority effective_chunksize = self._chunksize @@ -529,7 +542,50 @@ def _read_csv(self) -> TextFileReader | DataFrame: ) csv_engine = self._get_csv_engine(length, effective_chunksize) - read_csv_kwargs = { + read_csv_kwargs = self._get_csv_read_options(csv_engine, effective_chunksize) + + try: + with ExitStack() as stack: + source: str | TextIOWrapper = self.output_location + binary_columns = self._configure_binary_csv_read(read_csv_kwargs, pd.read_csv) + if binary_columns: + storage_options = read_csv_kwargs.pop("storage_options", None) or {} + self._csv_stream = stack.enter_context( + self._open_binary_csv_stream(binary_columns, storage_options) + ) + source = self._csv_stream + result = pd.read_csv(source, **read_csv_kwargs) + if not isinstance(result, pd.DataFrame): + # The chunk iterator takes ownership of the stream. + stack.pop_all() + + # Log performance information for large files + if length > self.LARGE_FILE_THRESHOLD_BYTES: + mode = "chunked" if effective_chunksize else "full" + chunksize = f" with chunksize={effective_chunksize}" if effective_chunksize else "" + _logger.info( + f"Reading {length} bytes from S3 in {mode} mode " + f"using {csv_engine} engine{chunksize}" + ) + + return result + + except Exception as e: + _logger.exception(f"Failed to read {self.output_location}.") + raise OperationalError(*e.args) from e + + def _get_csv_read_options(self, csv_engine: str, chunksize: int | None) -> dict[str, Any]: + """Build pandas options for Athena CSV or tab-separated results.""" + if self.output_location and self.output_location.endswith(".txt"): + sep = "\t" + header = None + names = [d[0] for d in self.description or []] + else: + sep = "," + header = 0 + names = None + + read_csv_kwargs: dict[str, Any] = { "sep": sep, "header": header, "names": names, @@ -546,7 +602,7 @@ def _read_csv(self) -> TextFileReader | DataFrame: "default_cache_type": self._cache_type, "max_workers": self._max_workers, }, - "chunksize": effective_chunksize, + "chunksize": chunksize, "engine": csv_engine, } @@ -558,23 +614,134 @@ def _read_csv(self) -> TextFileReader | DataFrame: read_csv_kwargs.update(self._kwargs) - try: - result = pd.read_csv(self.output_location, **read_csv_kwargs) - - # Log performance information for large files - if length > self.LARGE_FILE_THRESHOLD_BYTES: - mode = "chunked" if effective_chunksize else "full" - chunksize = f" with chunksize={effective_chunksize}" if effective_chunksize else "" - _logger.info( - f"Reading {length} bytes from S3 in {mode} mode " - f"using {csv_engine} engine{chunksize}" - ) + return read_csv_kwargs + + @staticmethod + def _resolve_csv_column_names( + column_names: list[Any], + read_csv_kwargs: dict[str, Any], + read_csv: Callable[..., DataFrame], + ) -> tuple[list[Any], set[Any]]: + """Resolve customized or duplicate column names with pandas' header parser.""" + header_buffer = StringIO() + csv.writer(header_buffer, quoting=csv.QUOTE_ALL).writerow(column_names) + header_options = { + key: read_csv_kwargs[key] + for key in ( + "sep", + "delimiter", + "names", + "engine", + "quoting", + "quotechar", + "doublequote", + "escapechar", + "skipinitialspace", + ) + if key in read_csv_kwargs + } + column_names = read_csv( + StringIO(header_buffer.getvalue()), header=0, nrows=0, **header_options + ).columns.tolist() + selected_names = set(column_names) + if read_csv_kwargs.get("usecols") is not None: + selected_names = set( + read_csv( + StringIO(header_buffer.getvalue()), + header=0, + nrows=0, + usecols=read_csv_kwargs["usecols"], + **header_options, + ).columns + ) + return column_names, selected_names + + def _can_preserve_binary_csv_nulls(self, read_csv_kwargs: dict[str, Any]) -> bool: + """Whether CSV settings support distinguishing binary NULL from empty values.""" + return not ( + "varbinary" not in self._converter.mappings + or "converters" in self._kwargs + or not self.output_location + or not self.output_location.endswith(".csv") + or read_csv_kwargs.get("header") != 0 + or read_csv_kwargs.get("skiprows") is not None + or read_csv_kwargs.get("dialect") is not None + or read_csv_kwargs.get("quoting") == csv.QUOTE_NONE + or read_csv_kwargs.get("quotechar", '"') != '"' + ) - return result + def _needs_csv_column_name_resolution(self, column_names: list[Any]) -> bool: + """Whether pandas must resolve column names instead of using Athena metadata.""" + return ( + len(set(column_names)) != len(column_names) + or not all(column_names) + or bool( + self._kwargs.keys() + & { + "names", + "usecols", + "sep", + "delimiter", + "doublequote", + "escapechar", + "skipinitialspace", + } + ) + ) - except Exception as e: - _logger.exception(f"Failed to read {self.output_location}.") - raise OperationalError(*e.args) from e + def _configure_binary_csv_read( + self, read_csv_kwargs: dict[str, Any], read_csv: Callable[..., DataFrame] + ) -> set[int]: + """Wrap binary converters and return column positions needing NULL preservation.""" + if not self._can_preserve_binary_csv_nulls(read_csv_kwargs): + return set() + + description = self.description or [] + binary_columns = {i for i, d in enumerate(description) if d[1] == "varbinary"} + if not binary_columns: + return set() + + column_names = [d[0] for d in description] + converters = read_csv_kwargs["converters"] + if self._needs_csv_column_name_resolution(column_names): + column_names, selected_names = self._resolve_csv_column_names( + column_names, read_csv_kwargs, read_csv + ) + if len(column_names) != len(description): + return set() + converters = { + name: self._converter.get(d[1]) + for name, d in zip(column_names, description, strict=True) + if d[1] in self._converter.mappings and name in selected_names + } + binary_columns = {i for i in binary_columns if column_names[i] in selected_names} + + if binary_columns: + for index in binary_columns: + name = column_names[index] + converters[name] = partial(_convert_binary_csv, converters[name]) + read_csv_kwargs["converters"] = converters + return binary_columns + + def _open_binary_csv_stream( + self, binary_columns: set[int], storage_options: dict[str, Any] + ) -> TextIOWrapper: + """Open a stream that preserves binary NULL fields and original CSV newlines.""" + with ExitStack() as stack: + source = stack.enter_context( + filesystem_open( + self.output_location, + mode="rt", + encoding="utf-8", + newline="", + **storage_options, + ) + ) + reader = stack.enter_context(BinaryCSVReader(source, binary_columns)) + buffer = stack.enter_context(BufferedReader(reader)) + stream = TextIOWrapper(buffer, encoding="utf-8", newline="") + stack.pop_all() + return stream def _read_parquet(self, engine) -> DataFrame: import pandas as pd @@ -695,6 +862,7 @@ def close(self) -> None: import pandas as pd super().close() + self._df_iter.close() self._df_iter = PandasDataFrameIterator(pd.DataFrame(), _no_trunc_date) self._iterrows = enumerate([]) self._data_manifest = [] diff --git a/pyathena/s3fs/reader.py b/pyathena/s3fs/reader.py index a48524d8..14de0101 100644 --- a/pyathena/s3fs/reader.py +++ b/pyathena/s3fs/reader.py @@ -120,9 +120,13 @@ def __next__(self) -> list[str | None]: Raises: StopIteration: When end of file is reached or reader is closed. """ + return self._parse_line(self._read_record().rstrip("\r\n")) + + def _read_record(self) -> str: + """Read a complete CSV record, retaining its original text and line endings.""" if self._file is None: raise StopIteration - line = self._file.readline() + line: str = self._file.readline() if not line: raise StopIteration @@ -138,7 +142,7 @@ def __next__(self) -> list[str | None]: # Only scan the new line, passing current quote state in_quotes = self._check_quote_state(next_line, in_quotes) - return self._parse_line(line.rstrip("\r\n")) + return line def _check_quote_state(self, text: str, starting_state: bool = False) -> bool: """Check quote state after processing text. @@ -150,17 +154,8 @@ def _check_quote_state(self, text: str, starting_state: bool = False) -> bool: Returns: True if we end inside an unclosed quote. """ - in_quotes = starting_state - i = 0 - while i < len(text): - if text[i] == '"': - if in_quotes and i + 1 < len(text) and text[i + 1] == '"': - # Escaped quote inside quoted field, skip both - i += 2 - continue - in_quotes = not in_quotes - i += 1 - return in_quotes + # Escaped quotes occur in pairs and do not change the quote state. + return starting_state != bool(text.count('"') % 2) def _parse_line(self, line: str) -> list[str | None]: """Parse a single CSV line preserving NULL vs empty string distinction. diff --git a/pyathena/sqlalchemy/base.py b/pyathena/sqlalchemy/base.py index b36589bb..2345b7b5 100644 --- a/pyathena/sqlalchemy/base.py +++ b/pyathena/sqlalchemy/base.py @@ -31,6 +31,7 @@ from pyathena.sqlalchemy.preparer import AthenaDMLIdentifierPreparer from pyathena.sqlalchemy.types import ( TINYINT, + AthenaBinary, AthenaDate, AthenaStruct, AthenaTimestamp, @@ -174,6 +175,9 @@ class AthenaDialect(DefaultDialect): ] colspecs: dict[type[Any], type[Any]] = { # noqa: RUF012 + types.LargeBinary: AthenaBinary, + types.BINARY: AthenaBinary, + types.VARBINARY: AthenaBinary, types.DATE: AthenaDate, types.DATETIME: AthenaTimestamp, types.TIMESTAMP: AthenaTimestamp, diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 6a119eb3..d0f1360a 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -294,7 +294,7 @@ def visit_cast(self, cast: Cast[Any], **kwargs): type_clause = "VARCHAR" elif isinstance(cast.type, types.CHAR) and cast.type.length is None: type_clause = "CHAR" - elif isinstance(cast.type, (types.BINARY, types.VARBINARY)): + elif isinstance(cast.type, (types.LargeBinary, types.BINARY, types.VARBINARY)): type_clause = "VARBINARY" elif hasattr(types, "DOUBLE") and isinstance(cast.type, types.DOUBLE): type_clause = "DOUBLE" diff --git a/pyathena/sqlalchemy/types.py b/pyathena/sqlalchemy/types.py index 906b6c16..b4b08cdc 100644 --- a/pyathena/sqlalchemy/types.py +++ b/pyathena/sqlalchemy/types.py @@ -27,6 +27,16 @@ def get_double_type() -> type[Any]: return types.FLOAT +class AthenaBinary(types.LargeBinary): + """SQLAlchemy binary type with Athena hexadecimal literals.""" + + def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[bytes]: + def process(value: bytes) -> str: + return f"X'{value.hex()}'" + + return process + + class AthenaTimestamp(TypeEngine[datetime]): """SQLAlchemy type for Athena TIMESTAMP values. diff --git a/tests/pyathena/aio/arrow/test_cursor.py b/tests/pyathena/aio/arrow/test_cursor.py index 5253dabe..78d4a716 100644 --- a/tests/pyathena/aio/arrow/test_cursor.py +++ b/tests/pyathena/aio/arrow/test_cursor.py @@ -7,6 +7,22 @@ class TestAioArrowCursor: + async def test_binary_null_vs_empty(self, aio_arrow_cursor): + query = """SELECT * FROM (VALUES + (1, CAST(NULL AS VARBINARY), 'null', CAST(NULL AS VARCHAR)), + (2, X'', 'empty', ''), + (3, X'00ff275c25', 'comma, quote" and' || chr(10) || 'newline', 'NULL') + ) AS t(id, value, label, text_value) ORDER BY id""" + await aio_arrow_cursor.execute(query) + result = aio_arrow_cursor + rows = await result.fetchall() + assert [row[:3] for row in rows] == [ + (1, None, "null"), + (2, b"", "empty"), + (3, b"\x00\xff'\\%", 'comma, quote" and\nnewline'), + ] + assert [row[3] for row in rows] == ["", "", "NULL"] + async def test_fetchone(self, aio_arrow_cursor): await aio_arrow_cursor.execute("SELECT * FROM one_row") assert aio_arrow_cursor.rownumber == 0 diff --git a/tests/pyathena/aio/pandas/test_cursor.py b/tests/pyathena/aio/pandas/test_cursor.py index 06b02fc8..ee51b1f9 100644 --- a/tests/pyathena/aio/pandas/test_cursor.py +++ b/tests/pyathena/aio/pandas/test_cursor.py @@ -7,6 +7,21 @@ class TestAioPandasCursor: + async def test_binary_null_vs_empty(self, aio_pandas_cursor): + query = """SELECT * FROM (VALUES + (1, CAST(NULL AS VARBINARY), 'null', CAST(NULL AS VARCHAR)), + (2, X'', 'empty', ''), + (3, X'00ff275c25', 'comma, quote" and' || chr(10) || 'newline', 'NULL') + ) AS t(id, value, label, text_value) ORDER BY id""" + await aio_pandas_cursor.execute(query, chunksize=2) + result = aio_pandas_cursor + rows = await result.fetchall() + assert [row[:3] for row in rows] == [ + (1, None, "null"), + (2, b"", "empty"), + (3, b"\x00\xff'\\%", 'comma, quote" and\nnewline'), + ] + async def test_fetchone(self, aio_pandas_cursor): await aio_pandas_cursor.execute("SELECT * FROM one_row") assert aio_pandas_cursor.rownumber == 0 diff --git a/tests/pyathena/aio/sqlalchemy/test_base.py b/tests/pyathena/aio/sqlalchemy/test_base.py index e65317e0..8f35cdea 100644 --- a/tests/pyathena/aio/sqlalchemy/test_base.py +++ b/tests/pyathena/aio/sqlalchemy/test_base.py @@ -1,12 +1,70 @@ import pytest import sqlalchemy -from sqlalchemy import text +from sqlalchemy import cast, literal, select, text, types from sqlalchemy.sql.schema import MetaData, Table from tests import ENV class TestAsyncSQLAlchemyAthena: + @pytest.mark.parametrize( + "async_engine", + [ + {"driver": driver} + for driver in ("aiorest", "aiopandas", "aioarrow", "aiopolars", "aios3fs") + ], + indirect=True, + ) + @pytest.mark.parametrize( + ("value", "expected"), + [ + (b"", b""), + (b"\x00\xff'\\%", b"\x00\xff'\\%"), + (bytes(range(256)), bytes(range(256))), + (bytearray(b"\x00\xff"), b"\x00\xff"), + (memoryview(b"\x00\xff"), b"\x00\xff"), + ], + ids=["empty", "special", "all_bytes", "bytearray", "memoryview"], + ) + async def test_binary_parameters_and_literals(self, async_engine, value, expected): + _, conn = async_engine + columns = [ + cast(literal(value, type_=type_, literal_execute=literal_execute), type_) + for type_ in (types.LargeBinary, types.BINARY, types.VARBINARY) + for literal_execute in (False, True) + ] + statement = select(*columns) + assert (await conn.execute(statement)).one() == (expected,) * len(columns) + compiled = statement.compile(dialect=conn.dialect, compile_kwargs={"literal_binds": True}) + assert (await conn.exec_driver_sql(str(compiled))).one() == (expected,) * len(columns) + + @pytest.mark.parametrize( + "async_engine", + [ + {"driver": "aiorest"}, + {"driver": "aiopandas"}, + {"driver": "aioarrow"}, + {"driver": "aiopolars"}, + {"driver": "aios3fs"}, + {"driver": "aiopandas", "unload": True}, + {"driver": "aioarrow", "unload": True}, + ], + indirect=["async_engine"], + ids=["rest", "pandas_csv", "arrow_csv", "polars", "s3fs", "pandas_unload", "arrow_unload"], + ) + async def test_binary_null_vs_empty(self, async_engine): + _, conn = async_engine + columns = [ + cast(literal(value, type_=type_, literal_execute=literal_execute), type_) + for type_ in (types.LargeBinary, types.BINARY, types.VARBINARY) + for value in (None, b"") + for literal_execute in (False, True) + ] + statement = select(*columns) + assert (await conn.execute(statement)).one() == (None, None, b"", b"") * 3 + compiled = statement.compile(dialect=conn.dialect, compile_kwargs={"literal_binds": True}) + assert (await conn.exec_driver_sql(str(compiled))).one() == (None, None, b"", b"") * 3 + @pytest.mark.parametrize( "async_engine", [ diff --git a/tests/pyathena/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index 85648a3f..d18a5be9 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -4,7 +4,7 @@ import pytest -from pyathena import ExecuteOptions +from pyathena import BINARY, Binary, ExecuteOptions from pyathena.aio.cursor import AioCursor from pyathena.error import DatabaseError, ProgrammingError from pyathena.model import AthenaQueryExecution @@ -15,6 +15,24 @@ class TestAioCursor: + @pytest.mark.parametrize( + ("value", "expected"), + [ + (b"", b""), + (b"\x00\xff'\\%", b"\x00\xff'\\%"), + (bytes(range(256)), bytes(range(256))), + (bytearray(b"\x00\xff"), b"\x00\xff"), + (memoryview(b"\x00\xff"), b"\x00\xff"), + (Binary(bytearray(b"abc")), b"abc"), + (None, None), + ], + ids=["empty", "special", "all_bytes", "bytearray", "memoryview", "dbapi_binary", "null"], + ) + async def test_binary_parameter(self, aio_cursor, value, expected): + await aio_cursor.execute("SELECT CAST(%(value)s AS VARBINARY)", {"value": value}) + assert await aio_cursor.fetchone() == (expected,) + assert aio_cursor.description[0][1] == BINARY + async def test_fetchone(self, aio_cursor): await aio_cursor.execute("SELECT * FROM one_row") assert aio_cursor.rowcount == -1 diff --git a/tests/pyathena/arrow/test_async_cursor.py b/tests/pyathena/arrow/test_async_cursor.py index a8fb770b..734961e3 100644 --- a/tests/pyathena/arrow/test_async_cursor.py +++ b/tests/pyathena/arrow/test_async_cursor.py @@ -16,6 +16,22 @@ class TestAsyncArrowCursor: + def test_binary_null_vs_empty(self, async_arrow_cursor): + query = """SELECT * FROM (VALUES + (1, CAST(NULL AS VARBINARY), 'null', CAST(NULL AS VARCHAR)), + (2, X'', 'empty', ''), + (3, X'00ff275c25', 'comma, quote" and' || chr(10) || 'newline', 'NULL') + ) AS t(id, value, label, text_value) ORDER BY id""" + _, future = async_arrow_cursor.execute(query) + result = future.result() + rows = result.fetchall() + assert [row[:3] for row in rows] == [ + (1, None, "null"), + (2, b"", "empty"), + (3, b"\x00\xff'\\%", 'comma, quote" and\nnewline'), + ] + assert [row[3] for row in rows] == ["", "", "NULL"] + @pytest.mark.parametrize( "async_arrow_cursor", [{"cursor_kwargs": {"unload": False}}, {"cursor_kwargs": {"unload": True}}], diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index 417d5c16..9b32854c 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -19,6 +19,26 @@ class TestArrowCursor: + def test_binary_null_vs_empty(self, arrow_cursor): + query = """SELECT * FROM (VALUES + (1, CAST(NULL AS VARBINARY), 'null', CAST(NULL AS VARCHAR)), + (2, X'', 'empty', ''), + (3, X'00ff275c25', 'comma, quote" and' || chr(10) || 'newline', 'NULL') + ) AS t(id, value, label, text_value) ORDER BY id""" + arrow_cursor.execute(query) + assert arrow_cursor.as_arrow().column("value").to_pylist() == [None, "", "00 ff 27 5c 25"] + rows = arrow_cursor.fetchall() + assert [row[:3] for row in rows] == [ + (1, None, "null"), + (2, b"", "empty"), + (3, b"\x00\xff'\\%", 'comma, quote" and\nnewline'), + ] + assert [row[3] for row in rows] == ["", "", "NULL"] + + def test_binary_single_null(self, arrow_cursor): + arrow_cursor.execute("SELECT CAST(NULL AS VARBINARY) AS value") + assert arrow_cursor.fetchall() == [(None,)] + @pytest.mark.parametrize( "arrow_cursor", [{"cursor_kwargs": {"unload": False}}, {"cursor_kwargs": {"unload": True}}], diff --git a/tests/pyathena/conftest.py b/tests/pyathena/conftest.py index c60dff67..7e0f249f 100644 --- a/tests/pyathena/conftest.py +++ b/tests/pyathena/conftest.py @@ -89,6 +89,7 @@ def create_engine(**kwargs): "row_format", "serdeproperties", "tblproperties", + "unload", "verify", ]: if arg in kwargs: @@ -107,6 +108,8 @@ def create_engine(**kwargs): def create_async_engine(**kwargs): driver = kwargs.pop("driver", "aiorest") conn_str = ASYNC_SQLALCHEMY_CONNECTION_STRING.replace("+aiorest", f"+{driver}") + if "unload" in kwargs: + conn_str += "&unload={unload}" return _create_async_engine( conn_str.format( region_name=ENV.region_name, diff --git a/tests/pyathena/pandas/test_async_cursor.py b/tests/pyathena/pandas/test_async_cursor.py index 28aa0b95..d052d6a6 100644 --- a/tests/pyathena/pandas/test_async_cursor.py +++ b/tests/pyathena/pandas/test_async_cursor.py @@ -19,6 +19,21 @@ class TestAsyncPandasCursor: + def test_binary_null_vs_empty(self, async_pandas_cursor): + query = """SELECT * FROM (VALUES + (1, CAST(NULL AS VARBINARY), 'null', CAST(NULL AS VARCHAR)), + (2, X'', 'empty', ''), + (3, X'00ff275c25', 'comma, quote" and' || chr(10) || 'newline', 'NULL') + ) AS t(id, value, label, text_value) ORDER BY id""" + _, future = async_pandas_cursor.execute(query, chunksize=2) + result = future.result() + rows = result.fetchall() + assert [row[:3] for row in rows] == [ + (1, None, "null"), + (2, b"", "empty"), + (3, b"\x00\xff'\\%", 'comma, quote" and\nnewline'), + ] + @pytest.mark.parametrize( ("async_pandas_cursor", "parquet_engine", "chunksize"), [ diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index d5e6d51c..34aa0968 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -1,4 +1,5 @@ import contextlib +import csv import math import random import string @@ -13,6 +14,7 @@ import pytest from pyathena.error import DatabaseError, ProgrammingError +from pyathena.pandas.converter import DefaultPandasTypeConverter from pyathena.pandas.cursor import PandasCursor from pyathena.pandas.result_set import AthenaPandasResultSet, PandasDataFrameIterator from tests import ENV @@ -20,6 +22,195 @@ class TestPandasCursor: + @pytest.mark.parametrize( + ("engine", "chunksize"), [("auto", None), ("c", 2), ("python", 2), ("pyarrow", None)] + ) + def test_binary_null_vs_empty(self, pandas_cursor, engine, chunksize): + query = """SELECT * FROM (VALUES + (1, CAST(NULL AS VARBINARY), 'null', CAST(NULL AS VARCHAR)), + (2, X'', 'empty', ''), + (3, X'00ff275c25', + 'comma, quote" and' || chr(13) || chr(10) || 'newline', 'NULL') + ) AS t(id, value, label, text_value) ORDER BY id""" + pandas_cursor.execute(query, engine=engine, chunksize=chunksize) + rows = pandas_cursor.fetchall() + assert pandas_cursor.result_set._csv_stream.closed + assert [row[:3] for row in rows] == [ + (1, None, "null"), + (2, b"", "empty"), + (3, b"\x00\xff'\\%", 'comma, quote" and\r\nnewline'), + ] + + assert pd.isna(rows[0][3]) + assert pd.isna(rows[1][3]) + assert rows[2][3] == "NULL" + + @pytest.mark.parametrize("chunksize", [None, 2]) + def test_binary_as_pandas(self, pandas_cursor, chunksize): + pandas_cursor.execute( + "SELECT * FROM (VALUES (1, CAST(NULL AS VARBINARY)), (2, X''), (3, X'00ff')) " + "AS t(id, value) ORDER BY id", + chunksize=chunksize, + storage_options={"connection": pandas_cursor.connection, "default_cache_type": "none"}, + ) + result = pandas_cursor.as_pandas() + df = pd.concat(list(result), ignore_index=True) if chunksize else result + assert df["value"].tolist() == [None, b"", b"\x00\xff"] + assert pandas_cursor.result_set._csv_stream.closed + + @pytest.mark.parametrize("engine", ["c", "python"]) + @pytest.mark.parametrize( + "close_method", + ["iterator", "generator", "context", "exhaust", "nrows", "get_chunk", "cursor", "execute"], + ) + def test_binary_csv_stream_close(self, pandas_cursor, engine, close_method): + pandas_cursor.execute( + "SELECT X'00' AS value, rpad('x', 4096, 'x') AS padding " + "FROM UNNEST(sequence(1, 100)) AS t(id)", + engine=engine, + chunksize=1, + nrows=2 if close_method == "nrows" else None, + ) + stream = pandas_cursor.result_set._csv_stream + source = stream.buffer.raw._reader._file + dataframe_iterator = pandas_cursor.as_pandas() + chunks = pandas_cursor.iter_chunks() if close_method == "generator" else dataframe_iterator + assert next(chunks)["value"].tolist() == [b"\x00"] + assert not stream.closed + assert not source.closed + + if close_method in ("iterator", "generator"): + chunks.close() + elif close_method == "context": + with chunks: + assert next(chunks)["value"].tolist() == [b"\x00"] + elif close_method == "exhaust": + assert sum(len(chunk) for chunk in chunks) == 99 + elif close_method == "nrows": + assert sum(len(chunk) for chunk in chunks) == 1 + elif close_method == "get_chunk": + assert len(chunks.get_chunk(99)) == 99 + with pytest.raises(StopIteration): + chunks.get_chunk(1) + elif close_method == "cursor": + pandas_cursor.close() + else: + pandas_cursor.execute("SELECT 1") + + assert stream.closed + assert source.closed + chunks.close() + with pytest.raises(StopIteration): + next(dataframe_iterator) + with pytest.raises(StopIteration): + dataframe_iterator.get_chunk() + if close_method != "execute": + assert pandas_cursor.fetchone() is None + + @pytest.mark.parametrize("engine", ["c", "python"]) + @pytest.mark.parametrize("read_method", ["next", "get_chunk"]) + def test_binary_csv_stream_close_on_read_error(self, pandas_cursor, engine, read_method): + pandas_cursor.execute( + "SELECT X'00' AS value, rpad('x', 4096, 'x') AS padding " + "FROM UNNEST(sequence(1, 100)) AS t(id)", + engine=engine, + chunksize=1, + dtype={"padding": "int64"}, + ) + stream = pandas_cursor.result_set._csv_stream + source = stream.buffer.raw._reader._file + chunks = pandas_cursor.as_pandas() + assert not stream.closed + assert not source.closed + read_chunk = chunks.__next__ if read_method == "next" else chunks.get_chunk + with pytest.raises(ValueError, match=r"invalid literal|Unable to convert column"): + read_chunk() + assert stream.closed + assert source.closed + with pytest.raises(StopIteration): + next(chunks) + with pytest.raises(StopIteration): + chunks.get_chunk() + assert pandas_cursor.fetchone() is None + + def test_binary_converter_override(self, pandas_cursor): + pandas_cursor.execute( + "SELECT CAST(NULL AS VARBINARY) AS value, X'' AS empty_value", + converters={"value": bytes.fromhex, "empty_value": bytes.fromhex}, + ) + assert pandas_cursor.fetchone() == (b"", b"") + pandas_cursor.execute("SELECT X'00ff' AS value", converters={}) + assert pandas_cursor.fetchone() == ("00 ff",) + pandas_cursor.execute("SELECT X'00ff' AS value", converters=None) + assert pandas_cursor.fetchone() == ("00 ff",) + + @pytest.mark.parametrize("engine", ["c", "pyarrow"]) + def test_binary_without_converter(self, engine): + converter = DefaultPandasTypeConverter() + for type_ in list(converter.mappings): + converter.remove(type_) + with ( + connect(cursor_class=PandasCursor, converter=converter) as conn, + conn.cursor() as cursor, + ): + cursor.execute( + "SELECT X'00ff' AS value, %(padding)s AS padding", + {"padding": "x" * 200}, + engine=engine, + ) + assert cursor.fetchone() == ("00 ff", "x" * 200) + + @pytest.mark.parametrize("engine", ["c", "python"]) + @pytest.mark.parametrize( + ("read_options", "expected"), + [ + ({}, [None, b"", b"\x00\xff"]), + ({"names": ["null_value", "empty_value", "value"]}, [None, b"", b"\x00\xff"]), + ({"usecols": [1, 2]}, [b"", b"\x00\xff"]), + ({"names": [0, 1, 2], "usecols": [1, 2]}, [b"", b"\x00\xff"]), + ], + ids=["duplicate_names", "renamed", "selected", "integer_names_selected"], + ) + def test_binary_dataframe_column_names(self, pandas_cursor, engine, read_options, expected): + pandas_cursor.execute( + "SELECT CAST(NULL AS VARBINARY) AS value, X'' AS value, X'00ff' AS value", + engine=engine, + **read_options, + ) + assert pandas_cursor.as_pandas().iloc[0].tolist() == expected + + @pytest.mark.parametrize("read_options", [{"quoting": 3}, {"quotechar": "'"}]) + def test_binary_custom_quoting(self, pandas_cursor, read_options): + pandas_cursor.execute("SELECT X'00ff' AS value", **read_options) + assert pandas_cursor.as_pandas().iloc[0].tolist() == ['"00 ff"'] + + @pytest.mark.parametrize("engine", ["c", "python"]) + def test_binary_custom_dialect(self, pandas_cursor, engine): + dialect = csv.excel() + dialect.quoting = csv.QUOTE_NONE + pandas_cursor.execute( + "SELECT CAST(NULL AS VARBINARY) AS value, 'text' AS label", + engine=engine, + dialect=dialect, + ) + row = pandas_cursor.as_pandas().iloc[0].tolist() + assert pd.isna(row[0]) + assert row[1] == '"text"' + + def test_binary_duplicate_name_without_converter(self): + converter = DefaultPandasTypeConverter() + converter.remove("varbinary") + with ( + connect(cursor_class=PandasCursor, converter=converter) as conn, + conn.cursor() as cursor, + ): + cursor.execute("SELECT true AS value, X'00ff' AS value") + assert cursor.as_pandas().iloc[0].tolist() == [True, "00 ff"] + + def test_binary_single_null(self, pandas_cursor): + pandas_cursor.execute("SELECT CAST(NULL AS VARBINARY) AS value") + assert pandas_cursor.fetchall() == [(None,)] + @pytest.mark.parametrize( ("pandas_cursor", "parquet_engine", "chunksize"), [ diff --git a/tests/pyathena/pandas/test_reader.py b/tests/pyathena/pandas/test_reader.py new file mode 100644 index 00000000..87270627 --- /dev/null +++ b/tests/pyathena/pandas/test_reader.py @@ -0,0 +1,38 @@ +import csv +from io import BufferedReader, StringIO, TextIOWrapper + +import pytest + +from pyathena.pandas.reader import _BINARY_NULL, BinaryCSVReader + + +class TestBinaryCSVReader: + """Tests for preserving binary NULL fields in Athena CSV results.""" + + @pytest.mark.parametrize("buffer_size", [1, 7, 8192]) + def test_null_vs_empty_binary(self, buffer_size): + source = StringIO( + '"first","text","last"\n,"comma, quote"" and\r\nnewline",""\n"00 ff","日本語",\n' + ) + with TextIOWrapper( + BufferedReader(BinaryCSVReader(source, {0, 2}), buffer_size), newline="" + ) as stream: + assert list(csv.reader(stream)) == [ + ["first", "text", "last"], + [_BINARY_NULL, 'comma, quote" and\r\nnewline', ""], + ["00 ff", "日本語", _BINARY_NULL], + ] + assert source.closed + + def test_close_before_eof(self): + source = StringIO('"value"\n"00 ff"\n') + with BinaryCSVReader(source, {0}) as stream: + assert stream.read(1) == b'"' + assert source.closed + + def test_preserves_other_fields(self): + source = StringIO('"binary","text"\r\n,unquoted\r\n"",\r\n') + with BinaryCSVReader(source, {0}) as stream: + assert stream.read().decode() == ( + f'"binary","text"\r\n{_BINARY_NULL},unquoted\r\n"",\r\n' + ) diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 7c0952e9..4b0b2dac 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -1507,6 +1507,65 @@ def test_cast_as_varchar(self, engine): ).scalar() assert actual == "1" + @pytest.mark.parametrize( + "engine", + [{"driver": driver} for driver in ("rest", "pandas", "arrow", "polars", "s3fs")], + indirect=True, + ) + @pytest.mark.parametrize( + ("value", "expected"), + [ + (b"", b""), + (b"\x00\xff'\\%", b"\x00\xff'\\%"), + (bytes(range(256)), bytes(range(256))), + (bytearray(b"\x00\xff"), b"\x00\xff"), + (memoryview(b"\x00\xff"), b"\x00\xff"), + ], + ids=["empty", "special", "all_bytes", "bytearray", "memoryview"], + ) + def test_binary_parameters_and_literals(self, engine, value, expected): + _, conn = engine + columns = [ + expression.cast( + expression.literal(value, type_=type_, literal_execute=literal_execute), type_ + ) + for type_ in (types.LargeBinary, types.BINARY, types.VARBINARY) + for literal_execute in (False, True) + ] + statement = select(*columns) + assert conn.execute(statement).one() == (expected,) * len(columns) + compiled = statement.compile(dialect=conn.dialect, compile_kwargs={"literal_binds": True}) + assert conn.exec_driver_sql(str(compiled)).one() == (expected,) * len(columns) + + @pytest.mark.parametrize( + "engine", + [ + {"driver": "rest"}, + {"driver": "pandas"}, + {"driver": "arrow"}, + {"driver": "polars"}, + {"driver": "s3fs"}, + {"driver": "pandas", "unload": True}, + {"driver": "arrow", "unload": True}, + ], + indirect=["engine"], + ids=["rest", "pandas_csv", "arrow_csv", "polars", "s3fs", "pandas_unload", "arrow_unload"], + ) + def test_binary_null_vs_empty(self, engine): + _, conn = engine + columns = [ + expression.cast( + expression.literal(value, type_=type_, literal_execute=literal_execute), type_ + ) + for type_ in (types.LargeBinary, types.BINARY, types.VARBINARY) + for value in (None, b"") + for literal_execute in (False, True) + ] + statement = select(*columns) + assert conn.execute(statement).one() == (None, None, b"", b"") * 3 + compiled = statement.compile(dialect=conn.dialect, compile_kwargs={"literal_binds": True}) + assert conn.exec_driver_sql(str(compiled)).one() == (None, None, b"", b"") * 3 + def test_cast_as_binary(self, engine): engine, conn = engine one_row_complex = Table("one_row_complex", MetaData(schema=ENV.schema), autoload_with=conn) @@ -1514,10 +1573,12 @@ def test_cast_as_binary(self, engine): sqlalchemy.select( expression.cast(one_row_complex.c.col_string, types.BINARY), expression.cast(one_row_complex.c.col_varchar, types.VARBINARY), + expression.cast(one_row_complex.c.col_string, types.LargeBinary), ) ).one() assert actual[0] == b"a string" assert actual[1] == b"varchar" + assert actual[2] == b"a string" def test_create_table_with_partition(self, engine): engine, conn = engine diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index 5e49f6af..0b48b4a8 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -15,7 +15,18 @@ import pytest -from pyathena import BINARY, BOOLEAN, DATE, DATETIME, JSON, NUMBER, STRING, TIME, ExecuteOptions +from pyathena import ( + BINARY, + BOOLEAN, + DATE, + DATETIME, + JSON, + NUMBER, + STRING, + TIME, + Binary, + ExecuteOptions, +) from pyathena.converter import _to_array, _to_map, _to_struct from pyathena.cursor import Cursor from pyathena.error import DatabaseError, NotSupportedError, ProgrammingError @@ -392,6 +403,24 @@ def test_null_param(self, cursor): cursor.execute("SELECT %(param)s FROM one_row", {"param": None}) assert cursor.fetchall() == [(None,)] + @pytest.mark.parametrize( + ("value", "expected"), + [ + (b"", b""), + (b"\x00\xff'\\%", b"\x00\xff'\\%"), + (bytes(range(256)), bytes(range(256))), + (bytearray(b"\x00\xff"), b"\x00\xff"), + (memoryview(b"\x00\xff"), b"\x00\xff"), + (Binary(bytearray(b"abc")), b"abc"), + (None, None), + ], + ids=["empty", "special", "all_bytes", "bytearray", "memoryview", "dbapi_binary", "null"], + ) + def test_binary_parameter(self, cursor, value, expected): + cursor.execute("SELECT CAST(%(value)s AS VARBINARY)", {"value": value}) + assert cursor.fetchone() == (expected,) + assert cursor.description[0][1] == BINARY + def test_no_params(self, cursor): pytest.raises(DatabaseError, lambda: cursor.execute("SELECT %(param)s FROM one_row")) pytest.raises(KeyError, lambda: cursor.execute("SELECT %(param)s FROM one_row", {"a": 1})) diff --git a/tests/pyathena/test_formatter.py b/tests/pyathena/test_formatter.py index 87da041d..4709f637 100644 --- a/tests/pyathena/test_formatter.py +++ b/tests/pyathena/test_formatter.py @@ -275,6 +275,24 @@ def test_format_str(self, formatter): ) assert actual == expected + @pytest.mark.parametrize( + ("value", "expected"), + [ + (b"", "X''"), + (b"abc", "X'616263'"), + (b"\x00\xff'\\%", "X'00ff275c25'"), + (bytearray(b"\x00\xff"), "X'00ff'"), + (memoryview(b"\x00\xff"), "X'00ff'"), + ], + ) + def test_format_binary(self, formatter, value, expected): + assert formatter.format("SELECT %(value)s", {"value": value}) == f"SELECT {expected}" + + def test_format_binary_sequence(self, formatter): + assert formatter.format("SELECT X'00' IN %(values)s", {"values": [b"\x00", b"\xff"]}) == ( + "SELECT X'00' IN (X'00', X'ff')" + ) + def test_format_unicode(self, formatter): expected = textwrap.dedent( """ diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index dc67d575..5981b951 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -1,7 +1,10 @@ import pytest -from sqlalchemy import Integer, func, select, testing +from sqlalchemy import Integer, MetaData, func, select, testing, types +from sqlalchemy import Table as SATable +from sqlalchemy import testing as sa_testing from sqlalchemy.testing import eq_ from sqlalchemy.testing.suite import * # noqa: F403 +from sqlalchemy.testing.suite import BinaryTest as _BinaryTest from sqlalchemy.testing.suite import CTETest as _CTETest from sqlalchemy.testing.suite import FetchLimitOffsetTest as _FetchLimitOffsetTest from sqlalchemy.testing.suite import HasTableTest as _HasTableTest @@ -10,7 +13,6 @@ from sqlalchemy.testing.suite import SimpleUpdateDeleteTest as _SimpleUpdateDeleteTest from sqlalchemy.testing.suite import StringTest as _StringTest -del BinaryTest # noqa: F821 del ComponentReflectionTest # noqa: F821 del ComponentReflectionTestExtra # noqa: F821 del CompositeKeyReflectionTest # noqa: F821 @@ -29,6 +31,35 @@ del UuidTest # noqa: F821 +class BinaryTest(_BinaryTest): + @sa_testing.combinations(types.LargeBinary, types.BINARY, types.VARBINARY, argnames="datatype") + @sa_testing.combinations( + ("empty", b""), + ("special", b"\x00\xff'\\%"), + ("all_bytes", bytes(range(256))), + argnames="data", + id_="ia", + ) + def test_literal(self, literal_round_trip, datatype, data): + literal_round_trip(datatype, [data], [data]) + + def test_reflected_binary_roundtrip(self, connection): + binary_table = self.tables.binary_table + data = b"\x00\xff'\\%" + connection.execute(binary_table.insert(), {"id": 1, "binary_data": data}) + reflected = SATable( + binary_table.name, + MetaData(), + schema=binary_table.schema, + autoload_with=connection, + ) + assert isinstance(reflected.c.binary_data.type, types.BINARY) + row = connection.execute( + select(reflected.c.binary_data).where(reflected.c.binary_data == data) + ).one() + assert row == (data,) + + class CTETest(_CTETest): @classmethod def define_tables(cls, metadata): From 048498b62c58f22f784471e8b536f924ce4dc846 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:34:58 +0900 Subject: [PATCH 04/20] Backport #789: Skip AWS integration jobs for external fork pull requests (cherry picked from commit a03bc2b4ea5dddcbf6230aafda0a359dce12ef0e) Co-Authored-By: Claude Opus 5.5 --- .github/workflows/test-suite.yaml | 2 ++ 1 file changed, 2 insertions(+) diff --git a/.github/workflows/test-suite.yaml b/.github/workflows/test-suite.yaml index 535bf9a8..4edf52b6 100644 --- a/.github/workflows/test-suite.yaml +++ b/.github/workflows/test-suite.yaml @@ -9,6 +9,8 @@ on: 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: From 27331ec8b3955cb90217b21bd8e3176d6d59b03a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:35:19 +0900 Subject: [PATCH 05/20] Backport #774: Support native SQLAlchemy ARRAY types and typed round trips (cherry picked from commit 08e97577224a05e32653b6a3e34ff76f6e571e89) Conflict resolution: pyathena/sqlalchemy/base.py and tests/sqlalchemy/test_suite.py conflicted because #771 and #777 (not backported) changed the surrounding code. In base.py, applied #774's changes unchanged. The pyathena.sqlalchemy.util import gains _split_type_arguments, and the pyathena.util import keeps 3.x's names without #777's throttling helpers. get_columns keeps 3.x's body, which #774 does not change. In test_suite.py, added the helper types and NativeArrayTest verbatim, together with the imports they need. Master had partly imported these through #771: fixtures, Column and Table from sqlalchemy.testing.schema, and String. Co-Authored-By: Claude Opus 5.5 --- docs/sqlalchemy.md | 206 ++++---- pyathena/formatter.py | 28 +- pyathena/sqlalchemy/array.py | 355 +++++++++++++ pyathena/sqlalchemy/base.py | 47 +- pyathena/sqlalchemy/compiler.py | 291 ++++++++++- pyathena/sqlalchemy/map.py | 61 +++ pyathena/sqlalchemy/requirements.py | 2 +- pyathena/sqlalchemy/struct.py | 76 +++ pyathena/sqlalchemy/temporal.py | 79 +++ pyathena/sqlalchemy/types.py | 248 +-------- pyathena/sqlalchemy/util.py | 32 ++ tests/pyathena/sqlalchemy/test_array.py | 576 +++++++++++++++++++++ tests/pyathena/sqlalchemy/test_base.py | 44 +- tests/pyathena/sqlalchemy/test_compiler.py | 32 +- tests/pyathena/sqlalchemy/test_map.py | 44 ++ tests/pyathena/sqlalchemy/test_struct.py | 71 +++ tests/pyathena/sqlalchemy/test_temporal.py | 49 ++ tests/pyathena/sqlalchemy/test_types.py | 198 +------ tests/sqlalchemy/test_suite.py | 249 ++++++++- 19 files changed, 2131 insertions(+), 557 deletions(-) create mode 100644 pyathena/sqlalchemy/array.py create mode 100644 pyathena/sqlalchemy/map.py create mode 100644 pyathena/sqlalchemy/struct.py create mode 100644 pyathena/sqlalchemy/temporal.py create mode 100644 tests/pyathena/sqlalchemy/test_array.py create mode 100644 tests/pyathena/sqlalchemy/test_map.py create mode 100644 tests/pyathena/sqlalchemy/test_struct.py create mode 100644 tests/pyathena/sqlalchemy/test_temporal.py diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index b7be6fe7..9149bfa2 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -952,7 +952,60 @@ value = map_data['key1'] # Direct access ### ARRAY type support -PyAthena provides comprehensive support for Amazon Athena's ARRAY data types, enabling you to work with ordered collections of data in your Python applications. +PyAthena supports SQLAlchemy's `ARRAY` type and the dialect-specific `AthenaArray` type. +Both store native Athena arrays and return typed Python collections. +Reflected ARRAY columns use `AthenaArray` and preserve their element types, including nested arrays, maps, rows, decimal precision, and string length where Athena retains it. + +```python +from sqlalchemy import ARRAY, Column, Integer, MetaData, String, Table, select + +metadata = MetaData() +events = Table( + "events", + metadata, + Column("id", Integer), + Column("numbers", ARRAY(Integer)), + Column("labels", ARRAY(String, dimensions=2)), +) + +connection.execute( + events.insert(), + {"id": 1, "numbers": [1, None, 3], "labels": [["one", "two"], ["a,b", "001"]]}, +) +row = connection.execute(select(events.c.numbers, events.c.labels)).one() +assert row.numbers == [1, None, 3] +assert row.labels == [["one", "two"], ["a,b", "001"]] +``` + +Use `dimensions=N` for a fixed number of dimensions with the standard SQLAlchemy type. +Without it, the table column has one dimension. +The existing `AthenaArray(AthenaArray(Integer))` spelling also remains supported; do not combine nested ARRAY types with `dimensions`. +`AthenaArray()` still defaults to string elements, and the uppercase `pyathena.sqlalchemy.types.ARRAY` alias remains available. +Set `as_tuple=True` to return tuples at each array dimension instead of lists. + +Bound parameters, multiple parameter sets, and SQLAlchemy literals use Athena's `ARRAY[...]` constructors. +An empty list represents an empty array, `None` represents SQL NULL, and individual elements can also be NULL. +Ordinary DB API list and tuple parameters retain their existing `IN (...)` formatting. +For decimal binds, specify `Numeric(precision, scale)`; a precision-free `Numeric()` raises a compilation error because Athena's bare DECIMAL cast would round fractional values to scale zero. +Athena `Float` uses 32-bit REAL values; use `Double` for 64-bit floating-point elements. + +Typed SQLAlchemy SELECT expressions use a JSON transport projection to preserve nested values and strings containing commas, quotes, whitespace, or the word `null`. +Scalar leaves for supported Athena table types are decoded according to the declared type, preserving decimal precision, dates, timestamps, and binary values. +SQL predicates and intermediate subqueries still operate on native arrays. +Arrays with unknown (`NullType`) elements keep the cursor's native conversion instead of using typed transport. +For ordered typed ARRAY results, use SQLAlchemy column expressions. +Textual ORDER BY clauses may name selected columns (including comma-separated names and direction/null placement); other textual expressions raise a compilation error to prevent ordering serialized values or referring to columns outside their scope. +String label references such as `.order_by("id")` also resolve FROM-table columns using SQLAlchemy's normal rules. +For DISTINCT and compound queries, ordering expressions must refer to selected columns. +Ordered, DISTINCT, and compound ARRAY queries require explicit SELECT columns: use SQLAlchemy column expressions or `literal_column()` instead of `text()` projections, and `select(table)` instead of a wildcard. +Literal SQL expressions need an explicit label, for example `literal_column("cardinality(items)").label("size")`. +This avoids dropping unnamed columns or exposing internal ordering columns when the result is wrapped. +An outer `TypeDecorator` retains its result processor as well as native ARRAY ordering. +Raw `text()` queries and direct DB API queries retain the cursor's existing conversion behavior described below; they do not receive this projection automatically. + +Compared with earlier releases, reflected ARRAY columns are no longer reported as `String`. +ARRAY DDL now renders integer elements as `INT` and row elements as `STRUCT<...>`, which Athena requires for nested DDL types. +Code that inspects reflected types or compares compiled SQL strings should account for these changes. #### Basic Usage @@ -974,7 +1027,7 @@ This creates a table definition equivalent to: ```sql CREATE TABLE orders ( id INTEGER, - item_ids ARRAY, + item_ids ARRAY, tags ARRAY, categories ARRAY ) @@ -982,82 +1035,67 @@ CREATE TABLE orders ( #### Querying ARRAY data -PyAthena automatically converts ARRAY data between different formats: - -```python -from sqlalchemy import create_engine, select - -# Query ARRAY data using ARRAY constructor -result = connection.execute( - select().from_statement( - text("SELECT ARRAY[1, 2, 3, 4, 5] as item_ids") - ) -).fetchone() +Use `select()` with `ARRAY` or `AthenaArray` columns whose element types are known, either declared explicitly or reflected from Athena, to receive typed Python collections. +For example, the `events` table above returns `numbers` as a list of integers and `labels` as nested lists of strings. +PyAthena projects these result columns as JSON in the generated SQL and converts the returned values according to the column types. +You do not need to call `json.loads()` on these results. +This preserves strings such as `"a,b"`, `"001"`, and `"null"` as strings, including within nested arrays. -# Access ARRAY data as Python list -item_ids = result.item_ids # [1, 2, 3, 4, 5] -``` +#### Direct cursor and textual SQL results -#### Complex ARRAY operations +Direct `cursor.execute()` calls and untyped SQLAlchemy `text()` queries use the cursor's existing conversion behavior. +The examples in this section use a standard REST `cursor` created as in [Usage](usage.md), with its default converter and no result type hints or custom converters. +Other cursor implementations have their own conversion behavior. -For arrays containing complex data types: +The standard converter already converts simple ARRAY values to Python lists: ```python -# Arrays with STRUCT elements -result = connection.execute( - select().from_statement( - text("SELECT ARRAY[ROW('Alice', 25), ROW('Bob', 30)] as users") - ) -).fetchone() - -users = result.users # [{"0": "Alice", "1": 25}, {"0": "Bob", "1": 30}] - -# Using CAST AS JSON for complex ARRAY operations -result = connection.execute( - select().from_statement( - text("SELECT CAST(ARRAY[1, 2, 3] AS JSON) as data") - ) -).fetchone() - -# Parse JSON result -import json -if isinstance(result.data, str): - array_data = json.loads(result.data) # [1, 2, 3] -else: - array_data = result.data # Already converted to list +result = cursor.execute("SELECT ARRAY[1, 2, 3] AS numbers").fetchone() +numbers = result[0] # [1, 2, 3] ``` -#### Data format support - -PyAthena supports multiple ARRAY data formats: +This conversion predates the typed SQLAlchemy ARRAY support described above. +Typed SQLAlchemy ARRAY support does not change the behavior of direct cursor queries or untyped `text()` queries. -**Athena Native Format:** - -```python -# Input: '[1, 2, 3]' -# Output: [1, 2, 3] +The standard ARRAY converter first tries JSON parsing, then a limited parser for Athena's native text representation. +The following examples show its behavior for strings received as ARRAY values: -# Input: '[apple, banana, cherry]' -# Output: ["apple", "banana", "cherry"] -``` +| ARRAY text | Python result | +|---|---| +| `[1, 2, 3]` | `[1, 2, 3]` (`list`) | +| `[one, two]` | `["one", "two"]` (`list`) | +| `[[1, 2], [3, 4]]` | `[[1, 2], [3, 4]]` (nested `list`) | +| `[[one, two], [three]]` | `"[[one, two], [three]]"` (`str`) | +| `[a, b=1]` | `"[a, b=1]"` (`str`) | -**JSON Format:** +The numeric examples are valid JSON. +The unquoted nested string array is not valid JSON, and the native parser does not support nested arrays. +In the last example, `b=1` is outside a `{...}` ROW element, so the native parser rejects it. +For these unsupported representations, the converter returns the original string rather than dropping unparsed elements. +A string result therefore does not necessarily mean the stored ARRAY is invalid. +Calling `json.loads()` on that native text will not resolve these cases because it is not JSON. -```python -# Input: '[1, 2, 3]' -# Output: [1, 2, 3] +A list result alone does not guarantee that the original elements were preserved. +For example, the native text `[a,b, null]` becomes `["a", "b", None]`, which would be incorrect for an original array of two strings, `["a,b", "null"]`. +The native representation does not distinguish commas within strings from element separators, or the string `"null"` from a NULL element. +Use typed SQLAlchemy SELECT expressions for string elements or nested arrays, or explicitly request JSON when writing SQL directly, as shown below. -# Input: '["apple", "banana", "cherry"]' -# Output: ["apple", "banana", "cherry"] -``` +#### Requesting JSON in direct SQL -**Complex Nested Arrays:** +Use `CAST(... AS JSON)` to have Athena return JSON instead of its native ARRAY text representation. +With the standard REST cursor and default converter, JSON results are already decoded into Python values: ```python -# Input: '[{name=John, age=30}, {name=Jane, age=25}]' -# Output: [{"name": "John", "age": 30}, {"name": "Jane", "age": 25}] +result = cursor.execute( + "SELECT CAST(ARRAY[ARRAY['one', 'two'], ARRAY['three']] AS JSON) AS labels" +).fetchone() +labels = result[0] # [["one", "two"], ["three"]] ``` +The same applies to an untyped SQLAlchemy `text()` query using `awsathena+rest` with the default converter. +No additional `json.loads()` call is needed in these examples. +This path returns JSON-derived Python values; use typed SQLAlchemy ARRAY columns when you need conversion according to declared element types such as `Numeric` or `Date`. + #### Type definitions AthenaArray supports various item types: @@ -1067,59 +1105,27 @@ from pyathena.sqlalchemy.types import AthenaArray, AthenaStruct, AthenaMap # Simple arrays AthenaArray(String) # ARRAY -AthenaArray(Integer) # ARRAY +AthenaArray(Integer) # ARRAY # Arrays of complex types AthenaArray(AthenaStruct(...)) # ARRAY> AthenaArray(AthenaMap(...)) # ARRAY> # Nested arrays -AthenaArray(AthenaArray(Integer)) # ARRAY> +AthenaArray(AthenaArray(Integer)) # ARRAY> ``` #### Best practices -1. **Use appropriate item types** in AthenaArray definitions: +1. Declare the ARRAY element type and use typed SQLAlchemy SELECT expressions to preserve nested values and scalar types. +2. For direct cursor or untyped `text()` queries, request `CAST(... AS JSON)` to preserve string elements and nested arrays. +3. Handle SQL NULL and empty arrays before accessing an element: ```python - AthenaArray(Integer) # For numeric arrays - AthenaArray(String) # For string arrays - AthenaArray(AthenaStruct(...)) # For arrays of structs + # A result from a typed ARRAY SELECT or the JSON query above + first_label = labels[0] if labels else None ``` -2. **Use CAST AS JSON** for complex array operations: - - ```sql - SELECT CAST(complex_array AS JSON) FROM table_name - ``` - -3. **Handle NULL values** appropriately in your application logic: - - ```python - if result.array_column is not None: - # Process array data - first_item = result.array_column[0] if result.array_column else None - ``` - -#### Migration from RAW strings - -**Before (raw string handling):** - -```python -result = cursor.execute("SELECT array_column FROM table").fetchone() -raw_data = result[0] # "[1, 2, 3]" -import json -parsed_data = json.loads(raw_data) -``` - -**After (automatic conversion):** - -```python -result = cursor.execute("SELECT array_column FROM table").fetchone() -array_data = result[0] # [1, 2, 3] - automatically converted -first_item = array_data[0] # Direct access -``` - ### JSON type support PyAthena provides support for Amazon Athena's JSON data type, enabling you to work with JSON data in your SQLAlchemy applications. The JSON type is primarily used with Data Manipulation Language (DML) operations in Athena. diff --git a/pyathena/formatter.py b/pyathena/formatter.py index bb305f7f..f4f5caea 100644 --- a/pyathena/formatter.py +++ b/pyathena/formatter.py @@ -7,9 +7,10 @@ from abc import ABCMeta, abstractmethod from collections.abc import Callable from copy import deepcopy +from dataclasses import dataclass from datetime import date, datetime, timezone from decimal import Decimal -from typing import Any +from typing import Any, Literal from pyathena.error import ProgrammingError from pyathena.model import AthenaCompression, AthenaFileFormat @@ -17,6 +18,14 @@ _logger = logging.getLogger(__name__) +@dataclass(frozen=True) +class _ComplexParameter: + """Typed complex value supplied by the SQLAlchemy dialect.""" + + constructor: Literal["ARRAY", "MAP", "ROW", "JSON_PARSE"] + values: tuple[Any, ...] + + class Formatter(metaclass=ABCMeta): """Abstract base class for formatting Python values for SQL queries. @@ -288,7 +297,24 @@ def _format_decimal(formatter: Formatter, escaper: Callable[[str], str], val: An return f"DECIMAL {escaped}" +def _format_complex( + formatter: Formatter, escaper: Callable[[str], str], val: _ComplexParameter +) -> str: + items = [] + for value in val.values: + if isinstance(value, (bytes, bytearray)): + items.append(f"X'{value.hex()}'") + continue + processor = formatter.get(value) + if processor is None: + raise TypeError(f"{type(value)} is not defined formatter.") + items.append(str(processor(formatter, escaper, value))) + opening, closing = ("[", "]") if val.constructor == "ARRAY" else ("(", ")") + return f"{val.constructor}{opening}{', '.join(items)}{closing}" + + _DEFAULT_FORMATTERS: dict[type[Any], Callable[[Formatter, Callable[[str], str], Any], Any]] = { + _ComplexParameter: _format_complex, type(None): _format_none, date: _format_date, datetime: _format_datetime, diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py new file mode 100644 index 00000000..17a19a90 --- /dev/null +++ b/pyathena/sqlalchemy/array.py @@ -0,0 +1,355 @@ +"""Athena ARRAY types, JSON result projection, and nested value processing.""" + +from __future__ import annotations + +import json +from collections.abc import Mapping +from datetime import date, datetime +from decimal import Decimal +from typing import Any + +from sqlalchemy import cast, exc, types +from sqlalchemy.sql import sqltypes +from sqlalchemy.sql.elements import ColumnElement +from sqlalchemy.sql.type_api import TypeEngine +from sqlalchemy.sql.visitors import InternalTraversal + +from pyathena.formatter import _ComplexParameter +from pyathena.sqlalchemy.map import AthenaMap +from pyathena.sqlalchemy.struct import AthenaStruct +from pyathena.sqlalchemy.temporal import AthenaDate, AthenaTimestamp + + +class AthenaArray(sqltypes.ARRAY[Any]): + """SQLAlchemy type for Athena ARRAY complex type. + + ARRAY represents an ordered collection of elements of the same type. + + Args: + item_type: SQLAlchemy type for array elements. Defaults to String. + as_tuple: Return tuples instead of lists. Defaults to False. + dimensions: Fixed number of array dimensions. Defaults to one dimension. + zero_indexes: Translate zero-based SQLAlchemy indexes to one-based SQL indexes. + + Example: + >>> from sqlalchemy import Column, Table, MetaData, types + >>> from pyathena.sqlalchemy.types import AthenaArray + >>> metadata = MetaData() + >>> posts = Table('posts', metadata, + ... Column('tags', AthenaArray(types.String)) + ... ) + + See Also: + AWS Athena ARRAY Type: + https://docs.aws.amazon.com/athena/latest/ug/arrays.html + """ + + __visit_name__ = "array" + + def __init__( + self, + item_type: Any = None, + as_tuple: bool = False, + dimensions: int | None = None, + zero_indexes: bool = False, + ) -> None: + if dimensions is not None and ( + isinstance(dimensions, bool) or not isinstance(dimensions, int) or dimensions < 1 + ): + raise ValueError("ARRAY dimensions must be a positive integer.") + item_type = item_type() if isinstance(item_type, type) else item_type + if isinstance(item_type, sqltypes.ARRAY): + if dimensions is not None: + raise ValueError("Use either nested ARRAY types or dimensions, not both.") + # Preserve the public nested AthenaArray constructor and item_type. + super().__init__(sqltypes.String(), as_tuple, dimensions, zero_indexes) + self.item_type = item_type + else: + super().__init__(item_type or sqltypes.String(), as_tuple, dimensions, zero_indexes) + + def bind_expression(self, bindvalue): + """Cast a bound ARRAY value to its declared Athena element type.""" + # The cast also gives empty arrays and NULL-only arrays their element type. + return cast(bindvalue, self)._annotate({"_pyathena_array_bind": True}) + + def bind_processor(self, dialect): + """Return a processor that marks native ARRAY, MAP, and ROW parameters.""" + return _ArrayValueProcessor(self, dialect).bind + + def literal_processor(self, dialect): + """Return a processor that renders typed Athena array literals.""" + return _ArrayValueProcessor(self, dialect).literal + + def column_expression(self, colexpr): + """Project the outer ARRAY result as JSON while retaining its Python type.""" + return ( + colexpr + if _ArrayTypeInspector.has_unknown_element(self) + else _ArrayJSONProjection(colexpr, self) + ) + + def result_processor(self, dialect, coltype): + """Return a processor that restores the declared Python element types.""" + return _ArrayValueProcessor(self, dialect).result + + +class ARRAY(AthenaArray): + """Uppercase alias for AthenaArray type.""" + + __visit_name__ = "ARRAY" + + +class _ArrayJSONProjection(ColumnElement[Any]): + """SQL expression that serializes an outer SELECT's ARRAY column as JSON. + + SQLAlchemy calls ``AthenaArray.column_expression`` for result columns, so + predicates and intermediate SELECTs keep using native ARRAY values. The + Athena statement compiler renders this wrapper as a JSON envelope; the + ARRAY result processor then restores its declared Python element types. + + ``type`` keeps the original column type, including an outer TypeDecorator's + result processor. ``array_type`` describes the native ARRAY value that the + compiler must serialize. This object represents SQL, not fetched row data. + """ + + __visit_name__ = "athena_array_json_projection" + inherit_cache = True + _traverse_internals = [ # noqa: RUF012 + ("element", InternalTraversal.dp_clauseelement), + ("type", InternalTraversal.dp_type), + ("array_type", InternalTraversal.dp_type), + ] + + def __init__(self, element, type_): + self.element = element + self.type = element.type + self.array_type = type_ + + +class _ArrayTypeInspector: + """Interpret nested ARRAY element types for SQL compilation and value conversion. + + Type inspection is shared by the compiler and value processors. Resolving + a TypeDecorator uses the current dialect; dimensions and unknown elements + can be inspected without one. + """ + + def __init__(self, dialect: Any) -> None: + self.dialect = dialect + + @staticmethod + def item_type(type_: sqltypes.ARRAY[Any]) -> TypeEngine[Any]: + if type_.dimensions is not None and type_.dimensions > 1: + return AthenaArray( + type_.item_type, + as_tuple=type_.as_tuple, + dimensions=type_.dimensions - 1, + zero_indexes=type_.zero_indexes, + ) + return type_.item_type + + def decorator_impl(self, type_: types.TypeDecorator[Any]) -> TypeEngine[Any]: + if self.dialect.name in type_._variant_mapping: + return type_._variant_mapping[self.dialect.name] + implementation = type_.load_dialect_impl(self.dialect) + if isinstance(implementation, AthenaTimestamp): + return types.TIMESTAMP() + if isinstance(implementation, AthenaDate): + return types.DATE() + return implementation + + @staticmethod + def has_unknown_element(type_: TypeEngine[Any]) -> bool: + if isinstance(type_, sqltypes.ARRAY): + return _ArrayTypeInspector.has_unknown_element(type_.item_type) + if isinstance(type_, AthenaMap): + return _ArrayTypeInspector.has_unknown_element( + type_.key_type + ) or _ArrayTypeInspector.has_unknown_element(type_.value_type) + if isinstance(type_, AthenaStruct): + return any( + _ArrayTypeInspector.has_unknown_element(field) for field in type_.fields.values() + ) + return isinstance(type_, types.NullType) + + +class _ArrayValueProcessor: + """Convert one declared ARRAY type between Python values and Athena transport. + + SQLAlchemy constructs processors per type and dialect. Keep that context + here and share the recursive ARRAY/MAP/ROW traversal across bind parameters, + SQL literals, and fetched JSON results. + """ + + def __init__(self, array_type: AthenaArray, dialect: Any) -> None: + self.array_type = array_type + self.dialect = dialect + self._type_inspector = _ArrayTypeInspector(dialect) + + def bind(self, value: Any) -> Any: + return self._bind(value, self.array_type) + + def literal(self, value: Any) -> str: + return self._literal(value, self.array_type) + + def result(self, value: Any) -> Any: + if value is None: + return None + if isinstance(value, str): + try: + value = json.loads(value) + except json.JSONDecodeError: + # Textual SQL does not receive column_expression. Preserve the + # DBAPI's raw fallback when native nested data is ambiguous. + return value + if isinstance(value, dict) and "_pyathena_array" in value: + value = value["_pyathena_array"] + return self._decode(value, self.array_type, self.array_type.as_tuple) + + @staticmethod + def _complex_values(value: Any, type_: TypeEngine[Any]): + if isinstance(type_, sqltypes.ARRAY): + if not isinstance(value, (list, tuple)): + raise TypeError("ARRAY values must be lists or tuples.") + item_type = _ArrayTypeInspector.item_type(type_) + return "ARRAY", [(item, item_type) for item in value] + if isinstance(type_, AthenaMap): + if not isinstance(value, Mapping): + raise TypeError("MAP values must be mappings.") + return "MAP", [ + (list(value), AthenaArray(type_.key_type)), + (list(value.values()), AthenaArray(type_.value_type)), + ] + if isinstance(type_, AthenaStruct): + if isinstance(value, Mapping): + if set(value) != set(type_.fields): + raise ValueError("ROW value fields must match the declared fields.") + values = [value[name] for name in type_.fields] + elif isinstance(value, (list, tuple)) and len(value) == len(type_.fields): + values = list(value) + else: + raise TypeError("ROW values must match the declared fields.") + return "ROW", list(zip(values, type_.fields.values(), strict=True)) + return None + + def _bind(self, value: Any, type_: TypeEngine[Any]) -> Any: + if isinstance(type_, types.TypeDecorator): + if ( + self.dialect.name not in type_._variant_mapping + and type(type_).bind_processor is not types.TypeDecorator.bind_processor + ): + processor = type_.bind_processor(self.dialect) + return processor(value) if processor else value + if self.dialect.name not in type_._variant_mapping and type_._has_bind_processor: + value = type_.process_bind_param(value, self.dialect) + return self._bind(value, self._type_inspector.decorator_impl(type_)) + if value is None: + return None + complex_values = self._complex_values(value, type_) + if complex_values is not None: + constructor, items = complex_values + return _ComplexParameter( + constructor, tuple(self._bind(item, item_type) for item, item_type in items) + ) + if isinstance(type_, types.JSON): + serializer = self.dialect._json_serializer or json.dumps + return _ComplexParameter("JSON_PARSE", (serializer(value),)) + if isinstance(value, (list, tuple, Mapping)): + raise TypeError("ARRAY element shape does not match its declared type.") + if isinstance(type_, (types.LargeBinary, types.BINARY, types.VARBINARY)): + return bytes(value) + if isinstance(type_, (types.Date, types.DateTime)): + return value + processor = type_.dialect_impl(self.dialect).bind_processor(self.dialect) + return processor(value) if processor else value + + def _literal(self, value: Any, type_: TypeEngine[Any]) -> str: + if isinstance(type_, types.TypeDecorator): + if ( + self.dialect.name not in type_._variant_mapping + and type(type_).literal_processor is not types.TypeDecorator.literal_processor + ): + literal_override = type_.literal_processor(self.dialect) + if literal_override is not None: + return literal_override(value) + if self.dialect.name not in type_._variant_mapping: + if type_._has_literal_processor: + value = type_.process_literal_param(value, self.dialect) + elif type_._has_bind_processor: + value = type_.process_bind_param(value, self.dialect) + return self._literal(value, self._type_inspector.decorator_impl(type_)) + if value is None: + return "NULL" + complex_values = self._complex_values(value, type_) + if complex_values is not None: + constructor, items = complex_values + opening, closing = ("[", "]") if constructor == "ARRAY" else ("(", ")") + values = ", ".join(self._literal(item, item_type) for item, item_type in items) + return f"{constructor}{opening}{values}{closing}" + if isinstance(type_, types.JSON): + serializer = self.dialect._json_serializer or json.dumps + processor = types.String().literal_processor(self.dialect) + return f"JSON_PARSE({processor(serializer(value))})" + if isinstance(value, (list, tuple, Mapping)): + raise TypeError("ARRAY element shape does not match its declared type.") + if isinstance(type_, (types.LargeBinary, types.BINARY, types.VARBINARY)): + return f"X'{bytes(value).hex()}'" + if isinstance(type_, types.DateTime) and isinstance(value, datetime): + return AthenaTimestamp.process(value) + if isinstance(type_, types.Date) and isinstance(value, date): + return AthenaDate.process(value) + processor = type_.dialect_impl(self.dialect).literal_processor(self.dialect) + if processor is None: + raise exc.CompileError(f"No ARRAY element literal processor for {type_!r}.") + return str(processor(value)) + + def _decode(self, value: Any, type_: TypeEngine[Any], as_tuple: bool = False) -> Any: + if isinstance(type_, types.TypeDecorator): + value = self._decode(value, self._type_inspector.decorator_impl(type_), as_tuple) + if ( + self.dialect.name not in type_._variant_mapping + and type(type_).result_processor is not types.TypeDecorator.result_processor + ): + processor = type_.result_processor(self.dialect, None) + return processor(value) if processor else value + if self.dialect.name not in type_._variant_mapping and type_._has_result_processor: + return type_.process_result_value(value, self.dialect) + return value + if value is None: + return None + if isinstance(type_, sqltypes.ARRAY): + item_type = _ArrayTypeInspector.item_type(type_) + items = [self._decode(item, item_type, as_tuple) for item in value] + return tuple(items) if as_tuple else items + if isinstance(type_, AthenaMap): + map_items = value.items() if isinstance(value, dict) else value + return { + self._decode(key, type_.key_type): self._decode(item, type_.value_type, as_tuple) + for key, item in map_items + } + if isinstance(type_, AthenaStruct): + if not type_.fields: + return value + return { + name: self._decode(value[name], field_type, as_tuple) + for name, field_type in type_.fields.items() + } + if isinstance(type_, types.JSON): + return value + if isinstance(type_, types.Boolean): + return value if isinstance(value, bool) else value.lower() == "true" + if isinstance(type_, types.Integer): + return int(value) + if isinstance(type_, types.Numeric): + return Decimal(value) if type_.asdecimal else float(value) + if isinstance(type_, (types.DateTime, AthenaTimestamp)): + return value if isinstance(value, datetime) else datetime.fromisoformat(value) + if isinstance(type_, (types.Date, AthenaDate)): + return value if isinstance(value, date) else date.fromisoformat(value) + if isinstance(type_, (types.LargeBinary, types.BINARY, types.VARBINARY)): + return value if isinstance(value, bytes) else bytes.fromhex(value) + if isinstance(type_, types.String): + value = str(value) + processor = type_.dialect_impl(self.dialect).result_processor(self.dialect, None) + return processor(value) if processor else value + return value diff --git a/pyathena/sqlalchemy/base.py b/pyathena/sqlalchemy/base.py index 2345b7b5..1c2fce4b 100644 --- a/pyathena/sqlalchemy/base.py +++ b/pyathena/sqlalchemy/base.py @@ -31,13 +31,15 @@ from pyathena.sqlalchemy.preparer import AthenaDMLIdentifierPreparer from pyathena.sqlalchemy.types import ( TINYINT, + AthenaArray, AthenaBinary, AthenaDate, + AthenaMap, AthenaStruct, AthenaTimestamp, get_double_type, ) -from pyathena.sqlalchemy.util import _HashableDict +from pyathena.sqlalchemy.util import _HashableDict, _split_type_arguments from pyathena.util import strtobool if TYPE_CHECKING: @@ -75,7 +77,7 @@ "timestamp": types.TIMESTAMP, "binary": types.BINARY, "varbinary": types.BINARY, - "array": types.String, + "array": AthenaArray, "map": types.String, "struct": AthenaStruct, "row": AthenaStruct, @@ -178,6 +180,7 @@ class AthenaDialect(DefaultDialect): types.LargeBinary: AthenaBinary, types.BINARY: AthenaBinary, types.VARBINARY: AthenaBinary, + types.ARRAY: AthenaArray, types.DATE: AthenaDate, types.DATETIME: AthenaTimestamp, types.TIMESTAMP: AthenaTimestamp, @@ -380,7 +383,8 @@ def get_columns(self, connection: Connection, table_name: str, schema: str | Non ] return columns - def _get_column_type(self, type_: str): + def _get_column_type(self, type_: str, _nested: bool = False): + type_ = type_.strip() match = self._pattern_column_type.match(type_) if match: name = match.group(1).lower() @@ -389,6 +393,40 @@ def _get_column_type(self, type_: str): name = type_.lower() length = None + if name == "array": + try: + return AthenaArray(self._get_column_type(length, _nested=True) if length else None) + except (TypeError, ValueError): + util.warn(f"Did not recognize type '{type_}'") + return types.NullType() + if _nested and name == "map" and length: + key, value = _split_type_arguments(length) + return AthenaMap( + self._get_column_type(key, _nested=True), + self._get_column_type(value, _nested=True), + ) + if _nested and name in ("row", "struct") and length: + fields = [] + for field in _split_type_arguments(length): + pattern = ( + r'\s*("(?:[^"]|"")*"|`(?:[^`]|``)*`|[^:]+)\s*:\s*(.+)' + if name == "struct" + else r'\s*("(?:[^"]|"")*"|`(?:[^`]|``)*`|[^\s:]+)(?:\s*:\s*|\s+)(.+)' + ) + match = re.fullmatch( + pattern, + field, + ) + if match is None: + raise ValueError(f"Invalid ROW field: {field!r}") + field_name, field_type = match.groups() + field_name = field_name.strip() + if field_name[0] in ('"', "`"): + quote = field_name[0] + field_name = field_name[1:-1].replace(quote * 2, quote) + fields.append((field_name, self._get_column_type(field_type, _nested=True))) + return AthenaStruct(*fields) + if name in self.ischema_names: col_type = self.ischema_names[name] else: @@ -398,8 +436,7 @@ def _get_column_type(self, type_: str): args = [] if length: if col_type is types.DECIMAL: - precision, scale = length.split(",") - args = [int(precision), int(scale)] + args = [int(arg) for arg in length.split(",")] elif col_type is types.CHAR or col_type is types.VARCHAR: args = [int(length)] diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index d0f1360a..aa0bbc80 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -1,25 +1,42 @@ from __future__ import annotations +import re from collections.abc import Mapping from typing import TYPE_CHECKING, Any, cast -from sqlalchemy import exc, types, util +from sqlalchemy import exc, select, types, util +from sqlalchemy.sql import operators, visitors +from sqlalchemy.sql import util as sql_util from sqlalchemy.sql.compiler import ( DDLCompiler, GenericTypeCompiler, IdentifierPreparer, SQLCompiler, ) -from sqlalchemy.sql.elements import BindParameter, Cast +from sqlalchemy.sql.elements import ( + BindParameter, + Cast, + TextClause, + UnaryExpression, + _label_reference, + _textual_label_reference, +) from sqlalchemy.sql.schema import Column +from sqlalchemy.sql.selectable import CompoundSelect from pyathena.model import ( AthenaFileFormat, AthenaPartitionTransform, AthenaRowFormatSerde, ) +from pyathena.sqlalchemy.array import _ArrayTypeInspector from pyathena.sqlalchemy.preparer import AthenaDDLIdentifierPreparer -from pyathena.sqlalchemy.types import AthenaArray, AthenaMap, AthenaStruct, get_double_type +from pyathena.sqlalchemy.types import ( + AthenaMap, + AthenaStruct, + get_double_type, +) +from pyathena.sqlalchemy.util import _split_type_arguments if TYPE_CHECKING: from sqlalchemy import ( @@ -93,7 +110,7 @@ def visit_TINYINT(self, type_: types.Integer, **kw: Any) -> str: return "TINYINT" def visit_INTEGER(self, type_: types.Integer, **kw: Any) -> str: - return "INTEGER" + return "INT" if kw.get("_athena_array_ddl") else "INTEGER" def visit_SMALLINT(self, type_: types.SmallInteger, **kw: Any) -> str: return "SMALLINT" @@ -177,7 +194,16 @@ def visit_struct(self, type_, **kw): 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}") + preparer = ( + AthenaDDLIdentifierPreparer(self.dialect) + if kw.get("_athena_array_ddl") + else self.dialect.identifier_preparer + ) + name = preparer.quote(field_name) + separator = ":" if kw.get("_athena_array_ddl") else " " + field_specs.append(f"{name}{separator}{field_type_str}") + if kw.get("_athena_array_ddl"): + return f"STRUCT<{', '.join(field_specs)}>" return f"ROW({', '.join(field_specs)})" return "ROW()" return "ROW()" @@ -196,8 +222,9 @@ def visit_MAP(self, type_, **kw): return self.visit_map(type_, **kw) def visit_array(self, type_, **kw): - if isinstance(type_, AthenaArray): - item_type_str = self.process(type_.item_type, **kw) + if isinstance(type_, types.ARRAY): + kw["_athena_array_ddl"] = True + item_type_str = self.process(_ArrayTypeInspector.item_type(type_), **kw) return f"ARRAY<{item_type_str}>" return "ARRAY" @@ -225,9 +252,177 @@ class AthenaStatementCompiler(SQLCompiler): https://docs.aws.amazon.com/athena/latest/ug/ddl-sql-reference.html """ + @util.memoized_property + def _array_type_inspector(self): + return _ArrayTypeInspector(self.dialect) + def visit_char_length_func(self, fn: Function[Any], **kw: Any) -> str: return f"length{self.function_argspec(fn, **kw)}" + def translate_select_structure(self, select_stmt, **kw): + """Keep DISTINCT and ordering on native arrays before result serialization.""" + if ( + not self.stack + and not kw.get("asfrom") + and not select_stmt._annotations.get("_pyathena_array_result") + and (select_stmt._distinct or select_stmt._order_by_clauses) + and any(self._has_array_result(column) for column in select_stmt.selected_columns) + ): + return self._array_result_select(select_stmt) + return select_stmt + + def visit_compound_select(self, cs, asfrom=False, compound_index=None, **kw): + if ( + not self.stack + and not asfrom + and any(self._has_array_result(column) for column in cs.selected_columns) + ): + original_columns = list(cs.selected_columns) + rendered = self.process(self._array_result_select(cs), **kw) + self._result_columns = [ + entry._replace(objects=(*entry.objects, original)) + for entry, original in zip(self._result_columns, original_columns, strict=True) + ] + return rendered + return super().visit_compound_select(cs, asfrom=asfrom, compound_index=compound_index, **kw) + + def _has_array_result(self, column): + type_ = column.type.dialect_impl(self.dialect) + while isinstance(type_, types.TypeDecorator): + type_ = self._array_type_inspector.decorator_impl(type_) + return isinstance(type_, types.ARRAY) and not _ArrayTypeInspector.has_unknown_element(type_) + + def _array_result_select(self, statement): + if any( + isinstance(column, TextClause) + or (getattr(column, "is_literal", False) and column.name.rstrip().endswith("*")) + for column in statement._all_selected_columns + ): + raise exc.CompileError( + "Ordered, DISTINCT, and compound ARRAY results require explicit SELECT columns; " + "use SQLAlchemy column expressions or literal_column() instead of text(), " + "and select(table) instead of a wildcard" + ) + if any( + getattr(column, "is_literal", False) + and not re.fullmatch(r'(?:[^\W\d]\w*|"(?:[^"]|"")+")', column.name) + for column in statement._all_selected_columns + ): + raise exc.CompileError( + "Literal SQL expressions in ordered, DISTINCT, and compound ARRAY results " + "require an explicit label; use literal_column(...).label(...)" + ) + columns = list(statement.selected_columns) + inner = statement.order_by(None).limit(None).offset(None) + ordering = [] + hidden: list[Any] = [] + label_resolve = ( + dict(statement.selected_columns.items()) + if isinstance(statement, CompoundSelect) + else statement._compile_state_factory(statement, self)._label_resolve_dict[0] + ) + + def resolve_label(element: Any, **kw: Any) -> Any: + if isinstance(element, _textual_label_reference): + try: + return label_resolve[element.element] + except KeyError as error: + raise exc.CompileError( + f"Can't resolve ARRAY ORDER BY label {element.element!r}" + ) from error + if isinstance(element, _label_reference): + return element.element + return None + + clauses: list[Any] = [] + for clause in statement._order_by_clauses: + if isinstance(clause, TextClause): + try: + clauses.extend(TextClause(part) for part in _split_type_arguments(clause.text)) + except ValueError as error: + raise exc.CompileError( + "Textual ARRAY ordering must name selected columns; " + "use SQLAlchemy column expressions for other ordering" + ) from error + else: + clauses.append(clause) + for clause in clauses: + if isinstance(clause, TextClause): + match = re.fullmatch( + r'\s*("(?:[^"]|"")+"|[\w]+)(?:\s+(ASC|DESC))?(?:\s+NULLS\s+(FIRST|LAST))?\s*', + clause.text, + re.IGNORECASE, + ) + if match: + name, direction, nulls = match.groups() + quoted = name.startswith('"') + name = name[1:-1].replace('""', '"') if quoted else name + target = None + if not quoted and name.isdigit() and 1 <= int(name) <= len(columns): + target = columns[int(name) - 1] + elif name in statement.selected_columns: + target = statement.selected_columns[name] + if target is not None: + clause = target + if direction: + clause = clause.desc() if direction.upper() == "DESC" else clause.asc() + if nulls: + clause = ( + clause.nulls_first() + if nulls.upper() == "FIRST" + else clause.nulls_last() + ) + if isinstance(clause, TextClause): + raise exc.CompileError( + "Textual ARRAY ordering must name selected columns; " + "use SQLAlchemy column expressions for other ordering" + ) + clause = visitors.replacement_traverse(clause, {}, resolve_label) + modifiers = [] + while isinstance(clause, UnaryExpression) and clause.modifier in ( + operators.asc_op, + operators.desc_op, + operators.nulls_first_op, + operators.nulls_last_op, + ): + modifiers.append(clause.modifier) + clause = clause.element + index = next((i for i, column in enumerate(columns) if column.compare(clause)), None) + if ( + index is None + and hasattr(inner, "add_columns") + and not inner._distinct + and not isinstance(clause, TextClause) + ): + name = f"_pyathena_order_{len(hidden)}" + while name in statement.selected_columns: + name += "_" + hidden.append(clause.label(name)) + index = len(columns) + len(hidden) - 1 + ordering.append((clause, index, modifiers)) + + if hidden: + inner = inner.add_columns(*hidden) + source = inner.subquery() + outer = select(*list(source.c)[: len(columns)]) + adapter = sql_util.ClauseAdapter(source) + for clause, index, modifiers in ordering: + expression = source.c[index] if index is not None else adapter.traverse(clause) + if any(from_ is not source for from_ in expression._from_objects): + raise exc.CompileError( + "DISTINCT and compound ARRAY ORDER BY expressions " + "must refer to selected columns" + ) + for modifier in reversed(modifiers): + expression = UnaryExpression(expression, modifier=modifier) + outer = outer.order_by(expression) + outer = outer.offset(statement._offset_clause) + if statement._fetch_clause is not None: + outer = outer.fetch(statement._fetch_clause, **statement._fetch_clause_options) + else: + outer = outer.limit(statement._limit_clause) + return outer._annotate({"_pyathena_array_result": True}) + def visit_filter_func(self, fn: Function[Any], **kw: Any) -> str: """Compile Athena filter() function with lambda expressions. @@ -288,6 +483,11 @@ def visit_truediv_binary(self, binary, operator, **kw): return super().visit_truediv_binary(binary, operator, **kw) def visit_cast(self, cast: Cast[Any], **kwargs): + if isinstance(cast.type, (types.ARRAY, AthenaMap, AthenaStruct)): + type_clause = self._complex_dml_type( + cast.type, implicit_bind=cast._annotations.get("_pyathena_array_bind", False) + ) + return f"CAST({self.process(cast.clause, **kwargs)} AS {type_clause})" if (isinstance(cast.type, types.VARCHAR) and cast.type.length is None) or isinstance( cast.type, types.String ): @@ -307,6 +507,83 @@ def visit_cast(self, cast: Cast[Any], **kwargs): type_clause = cast.typeclause._compiler_dispatch(self, **kwargs) return f"CAST({cast.clause._compiler_dispatch(self, **kwargs)} AS {type_clause})" + def _complex_dml_type(self, type_, *, implicit_bind=False): + if isinstance(type_, types.TypeDecorator): + return self._complex_dml_type( + self._array_type_inspector.decorator_impl(type_), implicit_bind=implicit_bind + ) + if isinstance(type_, types.NullType): + raise exc.CompileError("Bound ARRAY values require an explicit element type") + if isinstance(type_, types.ARRAY): + item = self._complex_dml_type( + _ArrayTypeInspector.item_type(type_), implicit_bind=implicit_bind + ) + return f"ARRAY({item})" + if isinstance(type_, AthenaMap): + return ( + f"MAP({self._complex_dml_type(type_.key_type, implicit_bind=implicit_bind)}, " + f"{self._complex_dml_type(type_.value_type, implicit_bind=implicit_bind)})" + ) + if isinstance(type_, AthenaStruct): + fields = ", ".join( + f"{self.preparer.quote(name)} " + f"{self._complex_dml_type(field_type, implicit_bind=implicit_bind)}" + for name, field_type in type_.fields.items() + ) + return f"ROW({fields})" + if isinstance(type_, types.String): + return "VARCHAR" + if isinstance(type_, (types.LargeBinary, types.BINARY, types.VARBINARY)): + return "VARBINARY" + if isinstance(type_, getattr(types, "Double", get_double_type())): + return "DOUBLE" + if isinstance(type_, types.Float): + return "REAL" + if implicit_bind and isinstance(type_, types.Numeric) and type_.precision is None: + raise exc.CompileError( + "ARRAY decimal binds require explicit Numeric precision; " + "specify precision and scale to avoid implicit rounding" + ) + return self.dialect.type_compiler_instance.process(type_) + + def visit_athena_array_json_projection(self, expression, **kw): + value = self.process(expression.element, **kw) + encoded = self._array_json(value, expression.array_type) + # An object envelope keeps SQL NULL and CSV null markers out of the transport. + return f"json_format(CAST(MAP(ARRAY['_pyathena_array'], ARRAY[{encoded}]) AS JSON))" + + def _array_json(self, value, type_, depth=0): + if isinstance(type_, types.TypeDecorator): + return self._array_json(value, self._array_type_inspector.decorator_impl(type_), depth) + # Each recursive value becomes JSON, including map keys and typed scalar leaves. + variable = f"_pyathena_array_{depth}" + if isinstance(type_, types.ARRAY): + child = self._array_json(variable, _ArrayTypeInspector.item_type(type_), depth + 1) + return f"CAST(transform({value}, {variable} -> {child}) AS JSON)" + if isinstance(type_, AthenaMap): + key = self._array_json(f"{variable}[1]", type_.key_type, depth + 1) + item = self._array_json(f"{variable}[2]", type_.value_type, depth + 1) + return ( + f"CAST(transform(map_entries({value}), {variable} -> ARRAY[{key}, {item}]) AS JSON)" + ) + if isinstance(type_, AthenaStruct) and type_.fields: + names = ", ".join( + self.render_literal_value(name, types.String()) for name in type_.fields + ) + fields = ", ".join( + self._array_json(f"({value}).{self.preparer.quote(name)}", field_type, depth + 1) + for name, field_type in type_.fields.items() + ) + return ( + f"IF({value} IS NULL, CAST(NULL AS JSON), " + f"CAST(MAP(ARRAY[{names}], ARRAY[{fields}]) AS JSON))" + ) + if isinstance(type_, (types.JSON, types.NullType, AthenaStruct)): + return f"CAST({value} AS JSON)" + if isinstance(type_, (types.LargeBinary, types.BINARY, types.VARBINARY)): + return f"CAST(to_hex({value}) AS JSON)" + return f"CAST(CAST({value} AS VARCHAR) AS JSON)" + def limit_clause(self, select: GenerativeSelect, **kw): text = [] if select._offset_clause is not None: diff --git a/pyathena/sqlalchemy/map.py b/pyathena/sqlalchemy/map.py new file mode 100644 index 00000000..02c8c84c --- /dev/null +++ b/pyathena/sqlalchemy/map.py @@ -0,0 +1,61 @@ +"""Athena MAP types.""" + +from __future__ import annotations + +from typing import Any + +from sqlalchemy.sql import sqltypes +from sqlalchemy.sql.type_api import TypeEngine + + +class AthenaMap(TypeEngine[dict[str, Any]]): + """SQLAlchemy type for Athena MAP complex type. + + MAP represents a collection of key-value pairs where all keys have the + same type and all values have the same type. + + Args: + key_type: SQLAlchemy type for map keys. Defaults to String. + value_type: SQLAlchemy type for map values. Defaults to String. + + Example: + >>> from sqlalchemy import Column, Table, MetaData, types + >>> from pyathena.sqlalchemy.types import AthenaMap + >>> metadata = MetaData() + >>> settings = Table('settings', metadata, + ... Column('config', AthenaMap(types.String, types.Integer)) + ... ) + + See Also: + AWS Athena MAP Type: + https://docs.aws.amazon.com/athena/latest/ug/maps.html + """ + + __visit_name__ = "map" + + def __init__(self, key_type: Any = None, value_type: Any = None) -> None: + if key_type is None: + self.key_type: TypeEngine[Any] = sqltypes.String() + elif isinstance(key_type, TypeEngine): + self.key_type = key_type + else: + # Assume it's a SQLAlchemy type class and instantiate it + self.key_type = key_type() + + if value_type is None: + self.value_type: TypeEngine[Any] = sqltypes.String() + elif isinstance(value_type, TypeEngine): + self.value_type = value_type + else: + # Assume it's a SQLAlchemy type class and instantiate it + self.value_type = value_type() + + @property + def python_type(self) -> type: + return dict + + +class MAP(AthenaMap): + """Uppercase alias for AthenaMap type.""" + + __visit_name__ = "MAP" diff --git a/pyathena/sqlalchemy/requirements.py b/pyathena/sqlalchemy/requirements.py index a9ccc0f4..12f22d53 100644 --- a/pyathena/sqlalchemy/requirements.py +++ b/pyathena/sqlalchemy/requirements.py @@ -8,7 +8,7 @@ class Requirements(SuiteRequirements): @property def array_type(self): - return unsupported() + return supported() @property def uuid_data_type(self): diff --git a/pyathena/sqlalchemy/struct.py b/pyathena/sqlalchemy/struct.py new file mode 100644 index 00000000..9df49974 --- /dev/null +++ b/pyathena/sqlalchemy/struct.py @@ -0,0 +1,76 @@ +"""Athena STRUCT/ROW types.""" + +from __future__ import annotations + +from typing import Any + +from sqlalchemy.sql import sqltypes +from sqlalchemy.sql.type_api import TypeEngine + + +class AthenaStruct(TypeEngine[dict[str, Any]]): + """SQLAlchemy type for Athena STRUCT/ROW complex type. + + STRUCT represents a record with named fields, similar to a database row + or a Python dictionary with typed values. Each field has a name and a + data type. + + Args: + *fields: Field specifications. Each can be either: + - A string (field name, defaults to STRING type) + - A tuple of (field_name, field_type) + + Example: + >>> from sqlalchemy import Column, Table, MetaData, types + >>> from pyathena.sqlalchemy.types import AthenaStruct + >>> metadata = MetaData() + >>> users = Table('users', metadata, + ... Column('address', AthenaStruct( + ... ('street', types.String), + ... ('city', types.String), + ... ('zip_code', types.Integer) + ... )) + ... ) + + See Also: + AWS Athena STRUCT Type: + https://docs.aws.amazon.com/athena/latest/ug/rows-and-structs.html + """ + + __visit_name__ = "struct" + + def __init__(self, *fields: str | tuple[str, Any]) -> None: + self.fields: dict[str, TypeEngine[Any]] = {} + + for field in fields: + if isinstance(field, str): + self.fields[field] = sqltypes.String() + elif isinstance(field, tuple) and len(field) == 2: + field_name, field_type = field + if isinstance(field_type, TypeEngine): + self.fields[field_name] = field_type + else: + # Assume it's a SQLAlchemy type class and instantiate it + self.fields[field_name] = field_type() + else: + raise ValueError(f"Invalid field specification: {field}") + + def __getitem__(self, key: str) -> TypeEngine[Any]: + return self.fields[key] + + @property + def _static_cache_key(self): + return ( + type(self), + tuple((name, type_._static_cache_key) for name, type_ in self.fields.items()), + ) + + @property + def python_type(self) -> type: + return dict + + +class STRUCT(AthenaStruct): + """Uppercase alias for AthenaStruct type.""" + + __visit_name__ = "STRUCT" diff --git a/pyathena/sqlalchemy/temporal.py b/pyathena/sqlalchemy/temporal.py new file mode 100644 index 00000000..f10c82b5 --- /dev/null +++ b/pyathena/sqlalchemy/temporal.py @@ -0,0 +1,79 @@ +"""Athena DATE and TIMESTAMP types and literal conversion.""" + +from __future__ import annotations + +from datetime import date, datetime +from typing import TYPE_CHECKING, Any + +from sqlalchemy.sql.type_api import TypeEngine + +if TYPE_CHECKING: + from sqlalchemy import Dialect + from sqlalchemy.sql.type_api import _LiteralProcessorType + + +class AthenaTimestamp(TypeEngine[datetime]): + """SQLAlchemy type for Athena TIMESTAMP values. + + This type handles the conversion of Python datetime objects to Athena's + TIMESTAMP literal syntax. When used in queries, datetime values are + rendered as ``TIMESTAMP 'YYYY-MM-DD HH:MM:SS.mmm'``. + + The type supports millisecond precision (3 decimal places) which matches + Athena's TIMESTAMP type precision. + + Example: + >>> from sqlalchemy import Column, Table, MetaData + >>> from pyathena.sqlalchemy.types import AthenaTimestamp + >>> metadata = MetaData() + >>> events = Table('events', metadata, + ... Column('event_time', AthenaTimestamp) + ... ) + """ + + __visit_name__ = "TIMESTAMP" + + render_literal_cast = True + render_bind_cast = True + + @staticmethod + def process(value: datetime | Any | None) -> str: + if isinstance(value, datetime): + return f"""TIMESTAMP '{value.strftime("%Y-%m-%d %H:%M:%S.%f")[:-3]}'""" + return f"TIMESTAMP '{value!s}'" + + def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[datetime] | None: + return self.process + + +class AthenaDate(TypeEngine[date]): + """SQLAlchemy type for Athena DATE values. + + This type handles the conversion of Python date objects to Athena's + DATE literal syntax. When used in queries, date values are rendered + as ``DATE 'YYYY-MM-DD'``. + + Example: + >>> from sqlalchemy import Column, Table, MetaData + >>> from pyathena.sqlalchemy.types import AthenaDate + >>> metadata = MetaData() + >>> orders = Table('orders', metadata, + ... Column('order_date', AthenaDate) + ... ) + """ + + __visit_name__ = "DATE" + + render_literal_cast = True + render_bind_cast = True + + @staticmethod + def process(value: date | Any) -> str: + # datetime is a subclass of date, so this branch also covers datetime, + # which is truncated to its date part. + if isinstance(value, date): + return f"DATE '{value:%Y-%m-%d}'" + return f"DATE '{value!s}'" + + def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[date] | None: + return self.process diff --git a/pyathena/sqlalchemy/types.py b/pyathena/sqlalchemy/types.py index b4b08cdc..e763717a 100644 --- a/pyathena/sqlalchemy/types.py +++ b/pyathena/sqlalchemy/types.py @@ -1,16 +1,40 @@ +"""Public SQLAlchemy type imports for PyAthena. + +Type-specific implementations live in their own modules. Re-export them here +so existing ``pyathena.sqlalchemy.types`` imports remain supported. +""" + from __future__ import annotations -from datetime import date, datetime from typing import TYPE_CHECKING, Any from sqlalchemy import types from sqlalchemy.sql import sqltypes -from sqlalchemy.sql.type_api import TypeEngine + +from pyathena.sqlalchemy.array import ARRAY, AthenaArray +from pyathena.sqlalchemy.map import MAP, AthenaMap +from pyathena.sqlalchemy.struct import STRUCT, AthenaStruct +from pyathena.sqlalchemy.temporal import AthenaDate, AthenaTimestamp if TYPE_CHECKING: from sqlalchemy import Dialect from sqlalchemy.sql.type_api import _LiteralProcessorType +__all__ = [ + "ARRAY", + "MAP", + "STRUCT", + "TINYINT", + "AthenaArray", + "AthenaBinary", + "AthenaDate", + "AthenaMap", + "AthenaStruct", + "AthenaTimestamp", + "Tinyint", + "get_double_type", +] + def get_double_type() -> type[Any]: """Get the appropriate type for DOUBLE based on SQLAlchemy version. @@ -37,69 +61,6 @@ def process(value: bytes) -> str: return process -class AthenaTimestamp(TypeEngine[datetime]): - """SQLAlchemy type for Athena TIMESTAMP values. - - This type handles the conversion of Python datetime objects to Athena's - TIMESTAMP literal syntax. When used in queries, datetime values are - rendered as ``TIMESTAMP 'YYYY-MM-DD HH:MM:SS.mmm'``. - - The type supports millisecond precision (3 decimal places) which matches - Athena's TIMESTAMP type precision. - - Example: - >>> from sqlalchemy import Column, Table, MetaData - >>> from pyathena.sqlalchemy.types import AthenaTimestamp - >>> metadata = MetaData() - >>> events = Table('events', metadata, - ... Column('event_time', AthenaTimestamp) - ... ) - """ - - render_literal_cast = True - render_bind_cast = True - - @staticmethod - def process(value: datetime | Any | None) -> str: - if isinstance(value, datetime): - return f"""TIMESTAMP '{value.strftime("%Y-%m-%d %H:%M:%S.%f")[:-3]}'""" - return f"TIMESTAMP '{value!s}'" - - def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[datetime] | None: - return self.process - - -class AthenaDate(TypeEngine[date]): - """SQLAlchemy type for Athena DATE values. - - This type handles the conversion of Python date objects to Athena's - DATE literal syntax. When used in queries, date values are rendered - as ``DATE 'YYYY-MM-DD'``. - - Example: - >>> from sqlalchemy import Column, Table, MetaData - >>> from pyathena.sqlalchemy.types import AthenaDate - >>> metadata = MetaData() - >>> orders = Table('orders', metadata, - ... Column('order_date', AthenaDate) - ... ) - """ - - render_literal_cast = True - render_bind_cast = True - - @staticmethod - def process(value: date | Any) -> str: - # datetime is a subclass of date, so this branch also covers datetime, - # which is truncated to its date part. - if isinstance(value, date): - return f"DATE '{value:%Y-%m-%d}'" - return f"DATE '{value!s}'" - - def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[date] | None: - return self.process - - class Tinyint(sqltypes.Integer): """SQLAlchemy type for Athena TINYINT (8-bit signed integer). @@ -117,160 +78,3 @@ class TINYINT(Tinyint): """ __visit_name__ = "TINYINT" - - -class AthenaStruct(TypeEngine[dict[str, Any]]): - """SQLAlchemy type for Athena STRUCT/ROW complex type. - - STRUCT represents a record with named fields, similar to a database row - or a Python dictionary with typed values. Each field has a name and a - data type. - - Args: - *fields: Field specifications. Each can be either: - - A string (field name, defaults to STRING type) - - A tuple of (field_name, field_type) - - Example: - >>> from sqlalchemy import Column, Table, MetaData, types - >>> from pyathena.sqlalchemy.types import AthenaStruct - >>> metadata = MetaData() - >>> users = Table('users', metadata, - ... Column('address', AthenaStruct( - ... ('street', types.String), - ... ('city', types.String), - ... ('zip_code', types.Integer) - ... )) - ... ) - - See Also: - AWS Athena STRUCT Type: - https://docs.aws.amazon.com/athena/latest/ug/rows-and-structs.html - """ - - __visit_name__ = "struct" - - def __init__(self, *fields: str | tuple[str, Any]) -> None: - self.fields: dict[str, TypeEngine[Any]] = {} - - for field in fields: - if isinstance(field, str): - self.fields[field] = sqltypes.String() - elif isinstance(field, tuple) and len(field) == 2: - field_name, field_type = field - if isinstance(field_type, TypeEngine): - self.fields[field_name] = field_type - else: - # Assume it's a SQLAlchemy type class and instantiate it - self.fields[field_name] = field_type() - else: - raise ValueError(f"Invalid field specification: {field}") - - def __getitem__(self, key: str) -> TypeEngine[Any]: - return self.fields[key] - - @property - def python_type(self) -> type: - return dict - - -class STRUCT(AthenaStruct): - """Uppercase alias for AthenaStruct type.""" - - __visit_name__ = "STRUCT" - - -class AthenaMap(TypeEngine[dict[str, Any]]): - """SQLAlchemy type for Athena MAP complex type. - - MAP represents a collection of key-value pairs where all keys have the - same type and all values have the same type. - - Args: - key_type: SQLAlchemy type for map keys. Defaults to String. - value_type: SQLAlchemy type for map values. Defaults to String. - - Example: - >>> from sqlalchemy import Column, Table, MetaData, types - >>> from pyathena.sqlalchemy.types import AthenaMap - >>> metadata = MetaData() - >>> settings = Table('settings', metadata, - ... Column('config', AthenaMap(types.String, types.Integer)) - ... ) - - See Also: - AWS Athena MAP Type: - https://docs.aws.amazon.com/athena/latest/ug/maps.html - """ - - __visit_name__ = "map" - - def __init__(self, key_type: Any = None, value_type: Any = None) -> None: - if key_type is None: - self.key_type: TypeEngine[Any] = sqltypes.String() - elif isinstance(key_type, TypeEngine): - self.key_type = key_type - else: - # Assume it's a SQLAlchemy type class and instantiate it - self.key_type = key_type() - - if value_type is None: - self.value_type: TypeEngine[Any] = sqltypes.String() - elif isinstance(value_type, TypeEngine): - self.value_type = value_type - else: - # Assume it's a SQLAlchemy type class and instantiate it - self.value_type = value_type() - - @property - def python_type(self) -> type: - return dict - - -class MAP(AthenaMap): - """Uppercase alias for AthenaMap type.""" - - __visit_name__ = "MAP" - - -class AthenaArray(TypeEngine[list[Any]]): - """SQLAlchemy type for Athena ARRAY complex type. - - ARRAY represents an ordered collection of elements of the same type. - - Args: - item_type: SQLAlchemy type for array elements. Defaults to String. - - Example: - >>> from sqlalchemy import Column, Table, MetaData, types - >>> from pyathena.sqlalchemy.types import AthenaArray - >>> metadata = MetaData() - >>> posts = Table('posts', metadata, - ... Column('tags', AthenaArray(types.String)) - ... ) - - See Also: - AWS Athena ARRAY Type: - https://docs.aws.amazon.com/athena/latest/ug/arrays.html - """ - - __visit_name__ = "array" - - def __init__(self, item_type: Any = None) -> None: - if item_type is None: - self.item_type: TypeEngine[Any] = sqltypes.String() - elif isinstance(item_type, TypeEngine): - self.item_type = item_type - else: - # Assume it's a SQLAlchemy type class and instantiate it - self.item_type = item_type() - - @property - def python_type(self) -> type: - return list - - -class ARRAY(AthenaArray): - """Uppercase alias for AthenaArray type.""" - - __visit_name__ = "ARRAY" diff --git a/pyathena/sqlalchemy/util.py b/pyathena/sqlalchemy/util.py index c82228d8..c55319c5 100644 --- a/pyathena/sqlalchemy/util.py +++ b/pyathena/sqlalchemy/util.py @@ -1,6 +1,38 @@ """Utility classes for PyAthena SQLAlchemy dialect.""" +def _split_type_arguments(value: str) -> list[str]: + """Split type arguments without splitting nested types or quoted field names.""" + parts = [] + start = 0 + brackets: list[str] = [] + quote: str | None = None + index = 0 + while index < len(value): + char = value[index] + if quote: + if char == quote: + if index + 1 < len(value) and value[index + 1] == quote: + index += 1 + else: + quote = None + elif char in ('"', "`"): + quote = char + elif char in "(<": + brackets.append(")" if char == "(" else ">") + elif char in ")>": + if not brackets or brackets.pop() != char: + raise ValueError(f"Unbalanced type arguments: {value!r}") + elif char == "," and not brackets: + parts.append(value[start:index].strip()) + start = index + 1 + index += 1 + if brackets or quote: + raise ValueError(f"Unbalanced type arguments: {value!r}") + parts.append(value[start:].strip()) + return parts + + class _HashableDict(dict): # type: ignore[type-arg] """A dictionary subclass that can be used as a dictionary key. diff --git a/tests/pyathena/sqlalchemy/test_array.py b/tests/pyathena/sqlalchemy/test_array.py new file mode 100644 index 00000000..5dce0576 --- /dev/null +++ b/tests/pyathena/sqlalchemy/test_array.py @@ -0,0 +1,576 @@ +import json +import pickle +from datetime import date, datetime +from decimal import Decimal +from enum import Enum +from types import SimpleNamespace + +import pytest +from sqlalchemy import ( + LABEL_STYLE_TABLENAME_PLUS_COL, + Column, + Integer, + MetaData, + String, + Table, + cast, + literal, + literal_column, + select, + text, + types, +) +from sqlalchemy import exc as sa_exc +from sqlalchemy.sql import sqltypes + +import pyathena +from pyathena.formatter import DefaultParameterFormatter +from pyathena.sqlalchemy.base import AthenaDialect +from pyathena.sqlalchemy.types import ( + ARRAY, + AthenaArray, + AthenaDate, + AthenaMap, + AthenaStruct, + AthenaTimestamp, +) + + +class Color(Enum): + RED = "red" + + +class OffsetInteger(types.TypeDecorator): + impl = Integer + cache_ok = True + + def process_bind_param(self, value, dialect): + return value - 1 if value is not None else None + + def process_result_value(self, value, dialect): + return value + 1 if value is not None else None + + +class DecoratedTimestamp(types.TypeDecorator): + impl = types.TIMESTAMP + cache_ok = True + + def process_result_value(self, value, dialect): + assert value is None or isinstance(value, datetime) + return value + + +class JSONEncodedDict(types.TypeDecorator): + impl = String + cache_ok = True + + def process_bind_param(self, value, dialect): + return json.dumps(value) + + def process_result_value(self, value, dialect): + return json.loads(value) + + +class TupleArray(types.TypeDecorator): + impl = AthenaArray(Integer) + cache_ok = True + + def process_result_value(self, value, dialect): + return tuple(value) if value is not None else None + + +class TestAthenaArray: + def test_creation_with_default(self): + array_type = AthenaArray() + assert isinstance(array_type.item_type, sqltypes.String) + + def test_creation_with_type_class(self): + array_type = AthenaArray(Integer) + assert isinstance(array_type.item_type, sqltypes.Integer) + + def test_creation_with_type_instance(self): + array_type = AthenaArray(Integer()) + assert isinstance(array_type.item_type, sqltypes.Integer) + + def test_creation_with_string_type(self): + array_type = AthenaArray(String) + assert isinstance(array_type.item_type, sqltypes.String) + + def test_python_type(self): + array_type = AthenaArray() + assert array_type.python_type is list + + def test_visit_name(self): + array_type = AthenaArray() + assert array_type.__visit_name__ == "array" + + def test_array_uppercase_visit_name(self): + array_type = ARRAY() + assert array_type.__visit_name__ == "ARRAY" + + def test_array_with_complex_type(self): + array_type = AthenaArray(AthenaStruct(("name", String), ("age", Integer))) + assert isinstance(array_type.item_type, AthenaStruct) + assert "name" in array_type.item_type.fields + assert "age" in array_type.item_type.fields + + def test_array_with_nested_array(self): + array_type = AthenaArray(AthenaArray(Integer)) + assert isinstance(array_type.item_type, AthenaArray) + assert isinstance(array_type.item_type.item_type, sqltypes.Integer) + + def test_array_with_map_type(self): + array_type = AthenaArray(AthenaMap(String, Integer)) + assert isinstance(array_type.item_type, AthenaMap) + assert isinstance(array_type.item_type.key_type, sqltypes.String) + assert isinstance(array_type.item_type.value_type, sqltypes.Integer) + + @pytest.mark.parametrize( + ("signature", "expected"), + [ + ("array", AthenaArray(types.INTEGER)), + ("ARRAY(ARRAY(VARCHAR(32)))", AthenaArray(AthenaArray(types.VARCHAR(32)))), + ("array", AthenaArray(types.DECIMAL(18, 7))), + ( + "array>>", + AthenaArray(AthenaMap(String, AthenaArray(types.INTEGER))), + ), + ( + 'array>', + AthenaArray( + AthenaStruct(("a,b", types.DECIMAL(10, 2)), ("c:d", types.VARCHAR(17))) + ), + ), + ], + ) + def test_array_reflection_preserves_element_types(self, signature, expected): + actual = AthenaDialect()._get_column_type(signature) + assert isinstance(actual, types.ARRAY) + assert actual._static_cache_key == expected._static_cache_key + + @pytest.mark.parametrize("dimensions", [0, -1, True, 1.5]) + def test_array_rejects_invalid_dimensions(self, dimensions): + with pytest.raises(ValueError, match="positive integer"): + AthenaArray(Integer, dimensions=dimensions) + + def test_array_rejects_ambiguous_dimensions(self): + with pytest.raises(ValueError, match="either nested ARRAY types or dimensions"): + AthenaArray(AthenaArray(Integer), dimensions=2) + + def test_array_insert_uses_typed_parameter(self): + formatter = DefaultParameterFormatter() + table = Table("array_values", MetaData(), Column("items", types.ARRAY(Integer))) + compiled = table.insert().values(items=[1, 2]).compile(dialect=AthenaDialect()) + params = { + name: compiled._bind_processors[name](value) for name, value in compiled.params.items() + } + assert ( + formatter.format(str(compiled), params) + == "INSERT INTO array_values (items) VALUES (CAST(ARRAY[1, 2] AS ARRAY(INTEGER)))" + ) + + def test_array_cache_key_includes_nested_fields(self): + first = AthenaArray(AthenaStruct(("x", Integer))) + second = AthenaArray(AthenaStruct(("x", String))) + assert first._static_cache_key != second._static_cache_key + assert hash(first._static_cache_key) + + @pytest.mark.parametrize( + ("item_type", "value"), + [ + (types.Numeric(), Decimal("1.50")), + (types.Numeric(scale=2), Decimal("1.50")), + (AthenaStruct(("amount", types.Numeric())), {"amount": Decimal("1.50")}), + ], + ) + def test_array_decimal_bind_requires_precision(self, item_type, value): + statement = select(literal([value], AthenaArray(item_type))) + with pytest.raises(sa_exc.CompileError, match="explicit Numeric precision"): + statement.compile(dialect=AthenaDialect()) + + def test_explicit_array_decimal_cast_keeps_default_precision(self): + value = literal([Decimal("1.50")], AthenaArray(types.Numeric(10, 2))) + sql = str(cast(value, AthenaArray(types.Numeric())).compile(dialect=AthenaDialect())) + assert sql == "CAST(CAST(%(param_1)s AS ARRAY(DECIMAL(10, 2))) AS ARRAY(DECIMAL))" + + @pytest.mark.parametrize( + "signature", + [ + "array>", + "array>", + "array", + "array(row(integer, varchar))", + ], + ) + def test_array_reflection_warns_for_unrecognized_nested_type(self, signature): + with pytest.warns(sa_exc.SAWarning, match="Did not recognize type"): + type_ = AthenaDialect()._get_column_type(signature) + assert isinstance(type_, types.NullType) + + +class TestArrayTypeInspector: + @pytest.mark.parametrize( + ("type_", "ddl", "dml"), + [ + (types.ARRAY(Integer), "ARRAY", "ARRAY(INTEGER)"), + (AthenaArray(), "ARRAY", "ARRAY(VARCHAR)"), + (ARRAY(String(12)), "ARRAY", "ARRAY(VARCHAR)"), + (types.ARRAY(String, dimensions=2), "ARRAY>", "ARRAY(ARRAY(VARCHAR))"), + (AthenaArray(AthenaArray(Integer)), "ARRAY>", "ARRAY(ARRAY(INTEGER))"), + (AthenaArray(types.Numeric(12, 3)), "ARRAY", "ARRAY(DECIMAL(12, 3))"), + (AthenaArray(types.Float), "ARRAY", "ARRAY(REAL)"), + (AthenaArray(types.BINARY), "ARRAY", "ARRAY(VARBINARY)"), + ], + ) + def test_array_type_rendering(self, type_, ddl, dml): + dialect = AthenaDialect() + assert type_.compile(dialect=dialect) == ddl + assert str(cast(literal(None), type_).compile(dialect=dialect)).endswith(f"AS {dml})") + assert isinstance(type_.dialect_impl(dialect), types.ARRAY) + + @pytest.mark.parametrize("item_type", [types.Double, types.DOUBLE, types.DOUBLE_PRECISION]) + def test_array_double_precision_cast(self, item_type): + sql = str(cast(literal(None), AthenaArray(item_type)).compile(dialect=AthenaDialect())) + assert "AS ARRAY(DOUBLE)" in sql + + def test_untyped_array_preserves_json_scalar_types(self): + dialect = AthenaDialect() + untyped = AthenaArray(types.NullType()) + sql = str(select(Column("items", untyped)).compile(dialect=dialect)) + assert "json_format" not in sql + assert untyped.result_processor(dialect, None)('[1,{"x":2},[3]]') == [1, {"x": 2}, [3]] + with pytest.raises(sa_exc.CompileError, match="explicit element type"): + select(literal([1], untyped)).compile(dialect=dialect) + + def test_unknown_array_does_not_rewrite_ordering(self): + table = Table("arrays", MetaData(), Column("items", AthenaArray(types.NullType()))) + sql = str(select(table).order_by(text("lower(name)")).compile(dialect=AthenaDialect())) + assert "ORDER BY lower(name)" in sql + assert "anon_1" not in sql + + @pytest.mark.parametrize( + ("item_type", "expected"), [(AthenaDate(), "DATE"), (AthenaTimestamp(), "TIMESTAMP")] + ) + def test_array_athena_temporal_element_type_compilation(self, item_type, expected): + dialect = AthenaDialect() + array = AthenaArray(item_type) + assert dialect.type_compiler_instance.process(array) == f"ARRAY<{expected}>" + assert f"AS ARRAY({expected})" in str(select(literal([], array)).compile(dialect=dialect)) + + +class TestArrayValueProcessor: + @pytest.mark.parametrize( + ("type_", "value", "expected"), + [ + (AthenaArray(Integer), [1, None, 3], "ARRAY[1, NULL, 3]"), + ( + AthenaArray(String), + ["thr'ee", "réve🐍 illé", "a,b", "null"], + "ARRAY['thr''ee', 'réve🐍 illé', 'a,b', 'null']", + ), + ( + AthenaArray(String, dimensions=2), + [["one"], [], None], + "ARRAY[ARRAY['one'], ARRAY[], NULL]", + ), + (AthenaArray(types.Date), [date(2025, 1, 2)], "ARRAY[DATE '2025-01-02']"), + (AthenaArray(AthenaDate), [date(2025, 1, 2)], "ARRAY[DATE '2025-01-02']"), + ( + AthenaArray(AthenaTimestamp), + [datetime(2025, 1, 2, 3, 4, 5)], + "ARRAY[TIMESTAMP '2025-01-02 03:04:05.000']", + ), + (AthenaArray(types.BINARY), [b"\x00\xff"], "ARRAY[X'00ff']"), + (AthenaArray(Integer), [], "ARRAY[]"), + (AthenaArray(Integer), None, "NULL"), + ], + ) + def test_array_bound_and_literal_values(self, type_, value, expected): + dialect = AthenaDialect() + assert type_.literal_processor(dialect)(value) == expected + bound = type_.bind_processor(dialect)(value) + actual = DefaultParameterFormatter().format("SELECT %(value)s", {"value": bound}) + assert actual == "SELECT " + expected.replace("NULL", "null") + + def test_array_binding_preserves_in_parameters(self): + formatter = DefaultParameterFormatter() + assert ( + formatter.format("SELECT 1 WHERE 1 IN %(items)s", {"items": [1, 2]}) + == "SELECT 1 WHERE 1 IN (1, 2)" + ) + + def test_array_binary_binding_uses_native_literals(self): + dialect = AthenaDialect(dbapi=pyathena) + processor = AthenaArray(types.BINARY).bind_processor(dialect) + assert ( + DefaultParameterFormatter().format("SELECT %(value)s", {"value": processor([b"\xff"])}) + == "SELECT ARRAY[X'ff']" + ) + + @pytest.mark.parametrize("value", [[[1]], "[1]", {"x": 1}]) + def test_array_binding_rejects_incorrect_shape(self, value): + with pytest.raises(TypeError, match="ARRAY"): + AthenaArray(Integer).bind_processor(AthenaDialect())(value) + + def test_array_textual_sql_preserves_native_fallback(self): + value = "[[one, two], [a,b]]" + assert ( + AthenaArray(String, dimensions=2).result_processor(AthenaDialect(), None)(value) + == value + ) + + @pytest.mark.parametrize( + ("type_", "encoded", "expected"), + [ + (AthenaArray(Integer), '["1",null,"3"]', [1, None, 3]), + (AthenaArray(String), '["001","null","a,b",""]', ["001", "null", "a,b", ""]), + ( + AthenaArray(Integer, dimensions=2, as_tuple=True), + '[["1"],[],null]', + ((1,), (), None), + ), + ( + AthenaArray(types.Numeric(30, 20)), + '["0.12345678901234567890"]', + [Decimal("0.12345678901234567890")], + ), + (AthenaArray(types.Date), '["2025-01-02"]', [date(2025, 1, 2)]), + (AthenaArray(AthenaDate), '["2025-01-02"]', [date(2025, 1, 2)]), + ( + AthenaArray(AthenaTimestamp), + '["2025-01-02 03:04:05"]', + [datetime(2025, 1, 2, 3, 4, 5)], + ), + (AthenaArray(types.BINARY), '["00FF",""]', [b"\x00\xff", b""]), + (AthenaArray(types.JSON), '[{"fraction":0.1}]', [{"fraction": 0.1}]), + ( + AthenaArray(AthenaMap(Integer, String)), + '[[["1","001"],["2",null]]]', + [{1: "001", 2: None}], + ), + ( + AthenaArray(AthenaStruct(("name", String), ("n", Integer))), + '[{"name":"001","n":"2"},null]', + [{"name": "001", "n": 2}, None], + ), + ], + ) + def test_array_result_conversion(self, type_, encoded, expected): + processor = type_.result_processor(AthenaDialect(), None) + assert processor(encoded) == expected + assert processor(json.loads(encoded)) == expected + assert processor(json.dumps({"_pyathena_array": json.loads(encoded)})) == expected + assert processor('{"_pyathena_array":null}') is None + assert processor(None) is None + + def test_array_custom_element_processors(self): + dialect = AthenaDialect() + enum = AthenaArray(types.Enum(Color)) + assert enum.result_processor(dialect, None)('["RED"]') == [Color.RED] + decorated = AthenaArray(OffsetInteger()) + assert decorated.result_processor(dialect, None)('["4"]') == [5] + sql = str(select(literal([5], decorated)).compile(dialect=dialect)) + assert "ARRAY(INTEGER)" in sql + + @pytest.mark.parametrize( + ("item_type", "value"), + [ + (DecoratedTimestamp(), datetime(2025, 1, 2, 3, 4, 5)), + (JSONEncodedDict(), {"x": 1}), + (OffsetInteger().with_variant(String(), "awsathena"), "unchanged"), + ], + ) + def test_decorator_bind_literal_and_result_paths(self, item_type, value): + dialect = AthenaDialect() + array = AthenaArray(item_type) + bind = array.bind_processor(dialect)([value]) + rendered = DefaultParameterFormatter().format("SELECT %(value)s", {"value": bind}) + assert "ARRAY[" in rendered + literal_sql = array.literal_processor(dialect)([value]) + assert "ARRAY[" in literal_sql + select(literal([value], array)).compile(dialect=dialect) + encoded = ( + value.isoformat(" ") + if isinstance(value, datetime) + else (json.dumps(value) if isinstance(value, dict) else value) + ) + assert array.result_processor(dialect, None)(json.dumps([encoded])) == [value] + + def test_array_pickle_type_uses_overridden_processors(self): + dialect = AthenaDialect(dbapi=SimpleNamespace(Binary=bytes, paramstyle="pyformat")) + array = AthenaArray(types.PickleType()) + bound = array.bind_processor(dialect)([5]) + assert pickle.loads(bound.values[0]) == 5 + assert array.result_processor(dialect, None)(json.dumps([bound.values[0].hex()])) == [5] + + +class TestArrayJSONProjection: + def test_array_result_projection_does_not_change_subquery_type(self): + table = Table("array_values", MetaData(), Column("items", AthenaArray(Integer))) + subquery = select(table.c["items"]).subquery() + compiled = str(select(subquery.c["items"]).compile(dialect=AthenaDialect())) + assert compiled.count("json_format(") == 1 + assert "SELECT array_values.items AS items" in compiled + assert isinstance(subquery.c["items"].type, types.ARRAY) + + def test_array_distinct_and_union_keep_native_ordering(self): + table = Table("arrays", MetaData(), Column("items", AthenaArray(Integer))) + for statement in ( + select(table.c["items"]).distinct().order_by("items"), + select(table.c["items"]).union_all(select(table.c["items"])).order_by("items"), + ): + sql = str(statement.compile(dialect=AthenaDialect())) + assert sql.count("json_format(") == 1 + assert "ORDER BY anon_1.items" in sql + + def test_array_textual_ordering_and_hive_field_spaces(self): + table = Table( + "arrays", MetaData(), Column("id", Integer), Column("items", AthenaArray(Integer)) + ) + sql = str(select(table).order_by(text("id DESC")).compile(dialect=AthenaDialect())) + assert "ORDER BY anon_1.id DESC" in sql + reflected = AthenaDialect()._get_column_type("array>") + assert list(reflected.item_type.fields) == ["first name"] + assert isinstance(reflected.item_type.fields["first name"], String) + + def test_textual_ordering_list_and_unresolved_expression(self): + table = Table( + "arrays", MetaData(), Column("id", Integer), Column("items", AthenaArray(Integer)) + ) + sql = str(select(table).order_by(text("items DESC, id")).compile(dialect=AthenaDialect())) + assert "ORDER BY anon_1.items DESC, anon_1.id" in sql + with pytest.raises(sa_exc.CompileError, match="column expressions"): + select(table).order_by(text("cardinality(items)")).compile(dialect=AthenaDialect()) + + def test_array_ordering_ordinals_use_selected_positions(self): + first = Table( + "first_table", MetaData(), Column("id", Integer), Column("items", AthenaArray(Integer)) + ) + second = Table("second_table", MetaData(), Column("id", Integer)) + sql = str( + select(first.c.id, second.c.id, first.c["items"]) + .order_by(text("2")) + .compile(dialect=AthenaDialect()) + ) + assert "ORDER BY anon_1.id_1" in sql + sql = str( + select(first) + .set_label_style(LABEL_STYLE_TABLENAME_PLUS_COL) + .order_by(text("1")) + .compile(dialect=AthenaDialect()) + ) + assert "ORDER BY anon_1.first_table_id" in sql + numeric = Table( + "numeric", MetaData(), Column("id", Integer), Column("1", AthenaArray(Integer)) + ) + sql = str(select(numeric).order_by(text('"1"')).compile(dialect=AthenaDialect())) + assert 'ORDER BY anon_1."1"' in sql + + @pytest.mark.parametrize( + "ordering", ["id > 5", "coalesce(name, ')')", "CASE WHEN id > 5 THEN 0 ELSE 1 END"] + ) + def test_unsupported_array_text_ordering_raises_compile_error(self, ordering): + table = Table("arrays", MetaData(), Column("items", AthenaArray(Integer))) + with pytest.raises(sa_exc.CompileError, match="column expressions"): + select(table).order_by(text(ordering)).compile(dialect=AthenaDialect()) + + def test_array_ordering_resolves_unselected_from_columns(self): + table = Table( + "arrays", MetaData(), Column("id", Integer), Column("items", AthenaArray(Integer)) + ) + statement = select(table.c["items"]).order_by("id") + sql = str(statement.compile(dialect=AthenaDialect())) + assert "arrays.id AS _pyathena_order_0" in sql + assert "ORDER BY anon_1._pyathena_order_0" in sql + + def test_array_ordering_resolves_qualified_selected_label(self): + table = Table("arrays", MetaData(), Column("items", AthenaArray(Integer))) + sql = str( + select(table.c["items"]).order_by("arrays_items").compile(dialect=AthenaDialect()) + ) + assert "ORDER BY anon_1.items" in sql + + def test_array_ordering_unknown_label_raises_compile_error(self): + table = Table("arrays", MetaData(), Column("items", AthenaArray(Integer))) + with pytest.raises(sa_exc.CompileError, match="resolve ARRAY ORDER BY label"): + select(table).order_by("missing").compile(dialect=AthenaDialect()) + + @pytest.mark.parametrize( + "projection", [text("id"), literal_column("*"), literal_column("arrays.*")] + ) + @pytest.mark.parametrize("operation", ["order_by", "distinct", "union_all"]) + def test_array_rewrite_rejects_untracked_projection(self, projection, operation): + table = Table( + "arrays", MetaData(), Column("id", Integer), Column("items", AthenaArray(Integer)) + ) + statement = select(projection, table.c["items"]) + if operation == "order_by": + statement = statement.order_by(table.c.id) + elif operation == "distinct": + statement = statement.distinct() + else: + statement = statement.union_all(statement) + with pytest.raises(sa_exc.CompileError, match="explicit SELECT columns"): + statement.compile(dialect=AthenaDialect()) + + def test_array_rewrite_keeps_explicit_literal_columns(self): + table = Table( + "arrays", MetaData(), Column("id", Integer), Column("items", AthenaArray(Integer)) + ) + sql = str( + select(literal_column("id"), table.c["items"]) + .order_by(table.c["items"]) + .compile(dialect=AthenaDialect()) + ) + assert sql.startswith("SELECT anon_1.id, json_format(") + + @pytest.mark.parametrize("compound", [False, True]) + def test_decorated_array_keeps_native_ordering_and_result_processor(self, compound): + dialect = AthenaDialect() + statement = select(literal([10], TupleArray()).label("items")) + if compound: + statement = statement.union_all(select(literal([2], TupleArray()).label("items"))) + compiled = statement.order_by("items").compile(dialect=dialect) + assert "ORDER BY anon_1.items" in str(compiled) + result_type = compiled._result_columns[0].type + assert isinstance(result_type, TupleArray) + processor = result_type.dialect_impl(dialect).result_processor(dialect, None) + assert processor('{"_pyathena_array":["2",null]}') == (2, None) + + def test_array_variant_keeps_transport_and_result_types(self): + dialect = AthenaDialect() + type_ = String().with_variant(AthenaArray(Integer), "awsathena") + compiled = ( + select(literal([2], type_).label("items")).order_by("items").compile(dialect=dialect) + ) + assert "transform(anon_1.items" in str(compiled) + processor = ( + compiled._result_columns[0].type.dialect_impl(dialect).result_processor(dialect, None) + ) + assert processor('{"_pyathena_array":["2"]}') == [2] + + @pytest.mark.parametrize("compound", [False, True]) + def test_array_ordering_rejects_columns_outside_distinct_or_union(self, compound): + table = Table( + "arrays", MetaData(), Column("id", Integer), Column("items", AthenaArray(Integer)) + ) + statement = select(table.c["items"]) + statement = statement.union_all(statement) if compound else statement.distinct() + with pytest.raises(sa_exc.CompileError, match="must refer to selected columns"): + statement.order_by(table.c.id).compile(dialect=AthenaDialect()) + sql = str(statement.order_by(table.c["items"]).compile(dialect=AthenaDialect())) + assert sql.count("FROM (") == 1 + assert "ORDER BY anon_1.items" in sql + + @pytest.mark.parametrize("expression", ["cardinality(items)", "arrays.id", "1"]) + def test_array_rewrite_requires_labels_for_literal_expressions(self, expression): + table = Table("arrays", MetaData(), Column("items", AthenaArray(Integer))) + value = literal_column(expression) + with pytest.raises(sa_exc.CompileError, match="require an explicit label"): + select(value, table.c["items"]).distinct().compile(dialect=AthenaDialect()) + sql = str( + select(value.label("value"), table.c["items"]) + .distinct() + .compile(dialect=AthenaDialect()) + ) + assert sql.startswith("SELECT anon_1.value, json_format(") diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 4b0b2dac..989ee191 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -47,6 +47,35 @@ def unique_s3tables_table_name(base: str) -> str: class TestSQLAlchemyAthena: + @pytest.mark.parametrize( + "engine", + [{"driver": driver} for driver in ("rest", "pandas", "arrow", "polars", "s3fs")], + indirect=True, + ) + def test_native_array_results_across_cursors(self, engine): + engine, _ = engine + modes = ( + (False, True) if engine.dialect.driver in ("pandas", "arrow", "polars") else (False,) + ) + for unload in modes: + url = engine.url.update_query_dict({"unload": str(unload).lower()}) + array_engine = sqlalchemy.create_engine(url) + try: + with array_engine.connect() as conn: + result = conn.execute( + select( + sqlalchemy.literal( + [["001", "a,b", "null", ""], [], None], + AthenaArray(types.String, dimensions=2), + ).label("nested"), + sqlalchemy.literal([], AthenaArray(types.Integer)).label("empty"), + sqlalchemy.literal(None, AthenaArray(types.Integer)).label("missing"), + ) + ).one() + assert tuple(result) == ([["001", "a,b", "null", ""], [], None], [], None) + finally: + array_engine.dispose() + @pytest.mark.parametrize( "engine", [ @@ -549,7 +578,8 @@ def test_reflect_select(self, engine): assert isinstance(one_row_complex.c.col_timestamp.type, types.TIMESTAMP) assert isinstance(one_row_complex.c.col_date.type, types.DATE) assert isinstance(one_row_complex.c.col_binary.type, types.BINARY) - assert isinstance(one_row_complex.c.col_array.type, types.String) + assert isinstance(one_row_complex.c.col_array.type, AthenaArray) + assert isinstance(one_row_complex.c.col_array.type.item_type, types.INTEGER) assert isinstance(one_row_complex.c.col_map.type, types.String) # With struct support, col_struct should now be recognized as AthenaStruct @@ -602,7 +632,7 @@ def test_get_column_type(self, engine): assert isinstance(dialect._get_column_type("timestamp"), types.TIMESTAMP) assert isinstance(dialect._get_column_type("date"), types.DATE) assert isinstance(dialect._get_column_type("binary"), types.BINARY) - assert isinstance(dialect._get_column_type("array"), types.String) + assert isinstance(dialect._get_column_type("array"), AthenaArray) assert isinstance(dialect._get_column_type("map"), types.String) # With struct support, struct types should be recognized as AthenaStruct @@ -2448,9 +2478,9 @@ def test_create_table_with_array_types(self, engine): # Verify ARRAY types are correctly compiled assert "tags ARRAY" in ddl_string - assert "scores 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.""" @@ -2527,7 +2557,7 @@ def test_create_table_with_struct_types(self, engine): "nested_struct ROW(personal ROW(first_name STRING, last_name STRING), " "preferences MAP)" in ddl_string ) - assert "struct_with_array ROW(tags ARRAY, scores ARRAY)" in ddl_string + assert "struct_with_array ROW(tags ARRAY, 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.""" @@ -2560,8 +2590,8 @@ 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 diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index 1fb120f5..16c9bbba 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -1,5 +1,3 @@ -from unittest.mock import Mock - import pytest from sqlalchemy import ( Column, @@ -27,14 +25,14 @@ class TestAthenaTypeCompiler: def test_visit_struct_empty(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = AthenaStruct() result = compiler.visit_struct(struct_type) assert result == "ROW()" def test_visit_struct_with_fields(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = AthenaStruct(("name", String), ("age", Integer)) result = compiler.visit_struct(struct_type) @@ -45,7 +43,7 @@ def test_visit_struct_with_fields(self): assert result.endswith(")") def test_visit_struct_uppercase(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = STRUCT(("id", Integer), ("title", String)) result = compiler.visit_STRUCT(struct_type) @@ -56,35 +54,35 @@ def test_visit_struct_uppercase(self): def test_visit_struct_no_fields_attribute(self): # Test struct type without fields attribute - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = type("MockStruct", (), {})() result = compiler.visit_struct(struct_type) assert result == "ROW()" def test_visit_struct_single_field(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) struct_type = AthenaStruct(("name", String)) result = compiler.visit_struct(struct_type) assert result == "ROW(name STRING)" or result == "ROW(name VARCHAR)" def test_visit_map_default(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) map_type = AthenaMap() result = compiler.visit_map(map_type) assert result == "MAP" def test_visit_map_with_types(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) map_type = AthenaMap(String, Integer) result = compiler.visit_map(map_type) assert result == "MAP" or result == "MAP" def test_visit_map_uppercase(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) map_type = MAP(Integer, String) result = compiler.visit_MAP(map_type) @@ -92,28 +90,28 @@ def test_visit_map_uppercase(self): def test_visit_map_no_attributes(self): # Test map type without key_type/value_type attributes - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) map_type = type("MockMap", (), {})() result = compiler.visit_map(map_type) assert result == "MAP" def test_visit_array_default(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) array_type = AthenaArray() result = compiler.visit_array(array_type) assert result == "ARRAY" def test_visit_array_with_type(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) array_type = AthenaArray(Integer) result = compiler.visit_array(array_type) - assert result == "ARRAY" + assert result == "ARRAY" def test_visit_array_uppercase(self): - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) array_type = ARRAY(String) result = compiler.visit_ARRAY(array_type) @@ -121,7 +119,7 @@ def test_visit_array_uppercase(self): def test_visit_array_no_attributes(self): # Test array type without item_type attribute - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) array_type = type("MockArray", (), {})() result = compiler.visit_array(array_type) @@ -131,7 +129,7 @@ def test_visit_json(self): """Test JSON type compilation.""" from sqlalchemy import types - dialect = Mock() + dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) json_type = types.JSON() result = compiler.visit_JSON(json_type) diff --git a/tests/pyathena/sqlalchemy/test_map.py b/tests/pyathena/sqlalchemy/test_map.py new file mode 100644 index 00000000..54f488d3 --- /dev/null +++ b/tests/pyathena/sqlalchemy/test_map.py @@ -0,0 +1,44 @@ +from sqlalchemy import ( + Integer, + String, +) +from sqlalchemy.sql import sqltypes + +from pyathena.sqlalchemy.types import ( + MAP, + AthenaMap, +) + + +class TestAthenaMap: + def test_creation_with_defaults(self): + map_type = AthenaMap() + assert isinstance(map_type.key_type, sqltypes.String) + assert isinstance(map_type.value_type, sqltypes.String) + + def test_creation_with_type_classes(self): + map_type = AthenaMap(String, Integer) + assert isinstance(map_type.key_type, sqltypes.String) + assert isinstance(map_type.value_type, sqltypes.Integer) + + def test_creation_with_type_instances(self): + map_type = AthenaMap(String(), Integer()) + assert isinstance(map_type.key_type, sqltypes.String) + assert isinstance(map_type.value_type, sqltypes.Integer) + + def test_python_type(self): + map_type = AthenaMap() + assert map_type.python_type is dict + + def test_visit_name(self): + map_type = AthenaMap() + assert map_type.__visit_name__ == "map" + + def test_map_uppercase_visit_name(self): + map_type = MAP() + assert map_type.__visit_name__ == "MAP" + + def test_mixed_type_definitions(self): + map_type = AthenaMap(String, Integer()) + assert isinstance(map_type.key_type, sqltypes.String) + assert isinstance(map_type.value_type, sqltypes.Integer) diff --git a/tests/pyathena/sqlalchemy/test_struct.py b/tests/pyathena/sqlalchemy/test_struct.py new file mode 100644 index 00000000..666e32ba --- /dev/null +++ b/tests/pyathena/sqlalchemy/test_struct.py @@ -0,0 +1,71 @@ +import pytest +from sqlalchemy import ( + Integer, + String, +) +from sqlalchemy.sql import sqltypes + +from pyathena.sqlalchemy.types import ( + STRUCT, + AthenaStruct, +) + + +class TestAthenaStruct: + def test_creation_with_strings(self): + struct_type = AthenaStruct("name", "age") + assert "name" in struct_type.fields + assert "age" in struct_type.fields + assert isinstance(struct_type.fields["name"], sqltypes.String) + assert isinstance(struct_type.fields["age"], sqltypes.String) + + def test_creation_with_tuples(self): + struct_type = AthenaStruct(("name", String), ("age", Integer)) + assert "name" in struct_type.fields + assert "age" in struct_type.fields + assert isinstance(struct_type.fields["name"], sqltypes.String) + assert isinstance(struct_type.fields["age"], sqltypes.Integer) + + def test_creation_with_type_instances(self): + struct_type = AthenaStruct(("name", String()), ("age", Integer())) + assert "name" in struct_type.fields + assert "age" in struct_type.fields + assert isinstance(struct_type.fields["name"], sqltypes.String) + assert isinstance(struct_type.fields["age"], sqltypes.Integer) + + def test_field_access_by_key(self): + struct_type = AthenaStruct(("name", String), ("age", Integer)) + name_field = struct_type["name"] + assert isinstance(name_field, sqltypes.String) + + def test_python_type(self): + struct_type = AthenaStruct(("name", String)) + assert struct_type.python_type is dict + + def test_invalid_field_specification(self): + with pytest.raises(ValueError, match="Invalid field specification"): + AthenaStruct(123) # Invalid field type + + def test_visit_name(self): + struct_type = AthenaStruct() + assert struct_type.__visit_name__ == "struct" + + def test_struct_uppercase_visit_name(self): + struct_type = STRUCT() + assert struct_type.__visit_name__ == "STRUCT" + + def test_empty_struct(self): + struct_type = AthenaStruct() + assert len(struct_type.fields) == 0 + + def test_mixed_field_definitions(self): + struct_type = AthenaStruct("name", ("age", Integer), ("active", String())) + assert len(struct_type.fields) == 3 + assert isinstance(struct_type.fields["name"], sqltypes.String) + assert isinstance(struct_type.fields["age"], sqltypes.Integer) + assert isinstance(struct_type.fields["active"], sqltypes.String) + + def test_field_access_nonexistent_key(self): + struct_type = AthenaStruct(("name", String)) + with pytest.raises(KeyError): + struct_type["nonexistent"] diff --git a/tests/pyathena/sqlalchemy/test_temporal.py b/tests/pyathena/sqlalchemy/test_temporal.py new file mode 100644 index 00000000..2fefba1c --- /dev/null +++ b/tests/pyathena/sqlalchemy/test_temporal.py @@ -0,0 +1,49 @@ +from datetime import date, datetime + +import pytest + +from pyathena.sqlalchemy.types import ( + AthenaDate, + AthenaTimestamp, +) + + +class TestAthenaDate: + @pytest.mark.parametrize( + ("value", "expected"), + [ + (date(2017, 1, 1), "DATE '2017-01-01'"), + (datetime(2017, 1, 1, 12, 34, 56), "DATE '2017-01-01'"), + ], + ) + def test_process_renders_date_only_literal(self, value, expected): + assert AthenaDate.process(value) == expected + + def test_process_falls_back_to_str(self): + assert AthenaDate.process("2017-01-01") == "DATE '2017-01-01'" + + +class TestAthenaTimestamp: + @pytest.mark.parametrize( + ("value", "expected"), + [ + # Athena TIMESTAMP has millisecond precision, so the six digits + # strftime("%f") emits are truncated to three. + ( + datetime(2017, 1, 1, 12, 34, 56, 789012), + "TIMESTAMP '2017-01-01 12:34:56.789'", + ), + ( + datetime(2017, 1, 1, 12, 34, 56), + "TIMESTAMP '2017-01-01 12:34:56.000'", + ), + ], + ) + def test_process_renders_millisecond_precision_literal(self, value, expected): + assert AthenaTimestamp.process(value) == expected + + def test_process_falls_back_to_str(self): + assert ( + AthenaTimestamp.process("2017-01-01 12:34:56.789") + == "TIMESTAMP '2017-01-01 12:34:56.789'" + ) diff --git a/tests/pyathena/sqlalchemy/test_types.py b/tests/pyathena/sqlalchemy/test_types.py index 3eb9c10c..d715e040 100644 --- a/tests/pyathena/sqlalchemy/test_types.py +++ b/tests/pyathena/sqlalchemy/test_types.py @@ -1,163 +1,12 @@ -from datetime import date, datetime - -import pytest -from sqlalchemy import Integer, String, types -from sqlalchemy.sql import sqltypes +from sqlalchemy import ( + types, +) from pyathena.sqlalchemy.types import ( - ARRAY, - MAP, - STRUCT, - AthenaArray, - AthenaDate, - AthenaMap, - AthenaStruct, - AthenaTimestamp, get_double_type, ) -class TestAthenaStruct: - def test_creation_with_strings(self): - struct_type = AthenaStruct("name", "age") - assert "name" in struct_type.fields - assert "age" in struct_type.fields - assert isinstance(struct_type.fields["name"], sqltypes.String) - assert isinstance(struct_type.fields["age"], sqltypes.String) - - def test_creation_with_tuples(self): - struct_type = AthenaStruct(("name", String), ("age", Integer)) - assert "name" in struct_type.fields - assert "age" in struct_type.fields - assert isinstance(struct_type.fields["name"], sqltypes.String) - assert isinstance(struct_type.fields["age"], sqltypes.Integer) - - def test_creation_with_type_instances(self): - struct_type = AthenaStruct(("name", String()), ("age", Integer())) - assert "name" in struct_type.fields - assert "age" in struct_type.fields - assert isinstance(struct_type.fields["name"], sqltypes.String) - assert isinstance(struct_type.fields["age"], sqltypes.Integer) - - def test_field_access_by_key(self): - struct_type = AthenaStruct(("name", String), ("age", Integer)) - name_field = struct_type["name"] - assert isinstance(name_field, sqltypes.String) - - def test_python_type(self): - struct_type = AthenaStruct(("name", String)) - assert struct_type.python_type is dict - - def test_invalid_field_specification(self): - with pytest.raises(ValueError, match="Invalid field specification"): - AthenaStruct(123) # Invalid field type - - def test_visit_name(self): - struct_type = AthenaStruct() - assert struct_type.__visit_name__ == "struct" - - def test_struct_uppercase_visit_name(self): - struct_type = STRUCT() - assert struct_type.__visit_name__ == "STRUCT" - - def test_empty_struct(self): - struct_type = AthenaStruct() - assert len(struct_type.fields) == 0 - - def test_mixed_field_definitions(self): - struct_type = AthenaStruct("name", ("age", Integer), ("active", String())) - assert len(struct_type.fields) == 3 - assert isinstance(struct_type.fields["name"], sqltypes.String) - assert isinstance(struct_type.fields["age"], sqltypes.Integer) - assert isinstance(struct_type.fields["active"], sqltypes.String) - - def test_field_access_nonexistent_key(self): - struct_type = AthenaStruct(("name", String)) - with pytest.raises(KeyError): - struct_type["nonexistent"] - - -class TestAthenaMap: - def test_creation_with_defaults(self): - map_type = AthenaMap() - assert isinstance(map_type.key_type, sqltypes.String) - assert isinstance(map_type.value_type, sqltypes.String) - - def test_creation_with_type_classes(self): - map_type = AthenaMap(String, Integer) - assert isinstance(map_type.key_type, sqltypes.String) - assert isinstance(map_type.value_type, sqltypes.Integer) - - def test_creation_with_type_instances(self): - map_type = AthenaMap(String(), Integer()) - assert isinstance(map_type.key_type, sqltypes.String) - assert isinstance(map_type.value_type, sqltypes.Integer) - - def test_python_type(self): - map_type = AthenaMap() - assert map_type.python_type is dict - - def test_visit_name(self): - map_type = AthenaMap() - assert map_type.__visit_name__ == "map" - - def test_map_uppercase_visit_name(self): - map_type = MAP() - assert map_type.__visit_name__ == "MAP" - - def test_mixed_type_definitions(self): - map_type = AthenaMap(String, Integer()) - assert isinstance(map_type.key_type, sqltypes.String) - assert isinstance(map_type.value_type, sqltypes.Integer) - - -class TestAthenaArray: - def test_creation_with_default(self): - array_type = AthenaArray() - assert isinstance(array_type.item_type, sqltypes.String) - - def test_creation_with_type_class(self): - array_type = AthenaArray(Integer) - assert isinstance(array_type.item_type, sqltypes.Integer) - - def test_creation_with_type_instance(self): - array_type = AthenaArray(Integer()) - assert isinstance(array_type.item_type, sqltypes.Integer) - - def test_creation_with_string_type(self): - array_type = AthenaArray(String) - assert isinstance(array_type.item_type, sqltypes.String) - - def test_python_type(self): - array_type = AthenaArray() - assert array_type.python_type is list - - def test_visit_name(self): - array_type = AthenaArray() - assert array_type.__visit_name__ == "array" - - def test_array_uppercase_visit_name(self): - array_type = ARRAY() - assert array_type.__visit_name__ == "ARRAY" - - def test_array_with_complex_type(self): - array_type = AthenaArray(AthenaStruct(("name", String), ("age", Integer))) - assert isinstance(array_type.item_type, AthenaStruct) - assert "name" in array_type.item_type.fields - assert "age" in array_type.item_type.fields - - def test_array_with_nested_array(self): - array_type = AthenaArray(AthenaArray(Integer)) - assert isinstance(array_type.item_type, AthenaArray) - assert isinstance(array_type.item_type.item_type, sqltypes.Integer) - - def test_array_with_map_type(self): - array_type = AthenaArray(AthenaMap(String, Integer)) - assert isinstance(array_type.item_type, AthenaMap) - assert isinstance(array_type.item_type.key_type, sqltypes.String) - assert isinstance(array_type.item_type.value_type, sqltypes.Integer) - - def test_get_double_type(): from pyathena.sqlalchemy.base import ischema_names @@ -167,44 +16,3 @@ def test_get_double_type(): else: assert result is types.FLOAT assert ischema_names["double"] is result - - -class TestAthenaDate: - @pytest.mark.parametrize( - ("value", "expected"), - [ - (date(2017, 1, 1), "DATE '2017-01-01'"), - (datetime(2017, 1, 1, 12, 34, 56), "DATE '2017-01-01'"), - ], - ) - def test_process_renders_date_only_literal(self, value, expected): - assert AthenaDate.process(value) == expected - - def test_process_falls_back_to_str(self): - assert AthenaDate.process("2017-01-01") == "DATE '2017-01-01'" - - -class TestAthenaTimestamp: - @pytest.mark.parametrize( - ("value", "expected"), - [ - # Athena TIMESTAMP has millisecond precision, so the six digits - # strftime("%f") emits are truncated to three. - ( - datetime(2017, 1, 1, 12, 34, 56, 789012), - "TIMESTAMP '2017-01-01 12:34:56.789'", - ), - ( - datetime(2017, 1, 1, 12, 34, 56), - "TIMESTAMP '2017-01-01 12:34:56.000'", - ), - ], - ) - def test_process_renders_millisecond_precision_literal(self, value, expected): - assert AthenaTimestamp.process(value) == expected - - def test_process_falls_back_to_str(self): - assert ( - AthenaTimestamp.process("2017-01-01 12:34:56.789") - == "TIMESTAMP '2017-01-01 12:34:56.789'" - ) diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 5981b951..f9c9d89b 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -1,8 +1,26 @@ +import json as _json +from datetime import date +from datetime import datetime as _datetime +from decimal import Decimal + import pytest -from sqlalchemy import Integer, MetaData, func, select, testing, types +from sqlalchemy import ( + Integer, + MetaData, + String, + cast, + func, + literal, + literal_column, + select, + testing, + text, + types, +) from sqlalchemy import Table as SATable from sqlalchemy import testing as sa_testing -from sqlalchemy.testing import eq_ +from sqlalchemy.testing import eq_, fixtures +from sqlalchemy.testing.schema import Column, Table from sqlalchemy.testing.suite import * # noqa: F403 from sqlalchemy.testing.suite import BinaryTest as _BinaryTest from sqlalchemy.testing.suite import CTETest as _CTETest @@ -13,6 +31,14 @@ from sqlalchemy.testing.suite import SimpleUpdateDeleteTest as _SimpleUpdateDeleteTest from sqlalchemy.testing.suite import StringTest as _StringTest +from pyathena.sqlalchemy.types import ( + AthenaArray, + AthenaDate, + AthenaMap, + AthenaStruct, + AthenaTimestamp, +) + del ComponentReflectionTest # noqa: F821 del ComponentReflectionTestExtra # noqa: F821 del CompositeKeyReflectionTest # noqa: F821 @@ -68,6 +94,225 @@ def define_tables(cls, metadata): metadata.tables["some_table"].c.parent_id.type = Integer() +class _ArrayTimestamp(types.TypeDecorator): + impl = types.TIMESTAMP + cache_ok = True + + def process_result_value(self, value, dialect): + assert value is None or isinstance(value, _datetime) + return value + + +class _ArrayJSONText(types.TypeDecorator): + impl = String + cache_ok = True + + def process_bind_param(self, value, dialect): + return _json.dumps(value) + + def process_result_value(self, value, dialect): + return _json.loads(value) + + +class _ArrayTuple(types.TypeDecorator): + impl = AthenaArray(Integer) + cache_ok = True + + def process_result_value(self, value, dialect): + return tuple(value) if value is not None else None + + +class NativeArrayTest(fixtures.TestBase): + __backend__ = True + __requires__ = ("array_type",) + + def test_native_ordering(self, connection, metadata): + table = Table( + "native_array_order", + metadata, + Column("id", Integer), + Column("value", AthenaArray(Integer)), + ) + table.create(connection) + connection.execute( + table.insert(), + [{"id": 1, "value": [10]}, {"id": 2, "value": [2]}, {"id": 3, "value": [2]}], + ) + value = table.c.value.label("items") + for ordering in (value, "items", text("items")): + stmt = select(value).distinct().order_by(ordering) + eq_(connection.execute(stmt).scalars().all(), [[2], [10]]) + eq_( + connection.execute(select(value).order_by(table.c.id.desc()).limit(2)).scalars().all(), + [[2], [2]], + ) + union = ( + select(table.c.value) + .where(table.c.id == 1) + .union_all(select(table.c.value).where(table.c.id == 2)) + .order_by("value") + ) + eq_(connection.execute(union).scalars().all(), [[2], [10]]) + eq_( + connection.execute(select(value, table.c.id).order_by(text("items DESC, id"))).all(), + [([10], 1), ([2], 2), ([2], 3)], + ) + eq_( + connection.execute(select(value).order_by("id")).scalars().all(), + [[10], [2], [2]], + ) + eq_( + connection.execute(select(table.c.value).order_by("native_array_order_value")) + .scalars() + .all(), + [[2], [2], [10]], + ) + + eq_( + connection.execute( + select(literal_column("cardinality(value)").label("size"), value).order_by( + table.c.id + ) + ).all(), + [(1, [10]), (1, [2]), (1, [2])], + ) + + def test_decorated_array_ordering(self, connection): + values = select(literal([10], _ArrayTuple()).label("items")).union_all( + select(literal([2], _ArrayTuple()).label("items")) + ) + eq_(connection.execute(values.order_by("items")).scalars().all(), [(2,), (10,)]) + source = values.subquery() + statement = select(source.c["items"]).distinct().order_by("items") + eq_(connection.execute(statement).scalars().all(), [(2,), (10,)]) + + def test_athena_temporal_elements(self, connection): + for item_type, value in ( + (AthenaDate(), date(2025, 1, 2)), + (AthenaTimestamp(), _datetime(2025, 1, 2, 3, 4, 5)), + ): + for literal_execute in (False, True): + statement = select( + literal([value], AthenaArray(item_type), literal_execute=literal_execute) + ) + eq_(connection.execute(statement).scalar_one(), [value]) + + def test_review_regressions(self, connection): + decimal_value = literal([Decimal("1.50")], AthenaArray(types.Numeric(10, 2))) + eq_( + connection.execute( + select(cast(decimal_value, AthenaArray(types.Numeric()))) + ).scalar_one(), + [Decimal("2")], + ) + expressions = [ + literal(["a-very-long-string"], AthenaArray(String(3))), + literal([0.1], AthenaArray(types.Double)), + literal([0.1], AthenaArray(types.DOUBLE_PRECISION)), + func.array_agg(func.length(literal("abc"))), + ] + # Aggregate and scalar expressions are checked separately for Athena grouping rules. + eq_( + tuple(connection.execute(select(*expressions[:3])).one()), + (["a-very-long-string"], [0.1], [0.1]), + ) + eq_(connection.execute(select(expressions[3])).scalar_one(), [3]) + custom_values = [ + (AthenaArray(_ArrayTimestamp()), [_datetime(2025, 1, 2, 3, 4, 5)]), + (AthenaArray(_ArrayJSONText()), [{"nested": [1, 2], "fraction": 0.1}]), + ] + for type_, value in custom_values: + for literal_execute in (False, True): + eq_( + connection.execute( + select(literal(value, type_, literal_execute=literal_execute)) + ).scalar_one(), + value, + ) + + def test_reflection_and_executemany(self, connection, metadata): + table = Table( + "native_array_values", + metadata, + Column("id", Integer), + Column("numbers", types.ARRAY(Integer)), + Column("labels", types.ARRAY(String, dimensions=2)), + Column("amounts", AthenaArray(types.Numeric(30, 20))), + ) + table.create(connection) + values = [ + { + "id": 1, + "numbers": [1, None, 3], + "labels": [["001", "null", "a,b", ""], ["thr'ee", "réve🐍 illé"]], + "amounts": [Decimal("0.12345678901234567890")], + }, + {"id": 2, "numbers": [], "labels": [[], None], "amounts": []}, + {"id": 3, "numbers": None, "labels": None, "amounts": None}, + ] + connection.execute(table.insert(), values) + reflected = Table(table.name, MetaData(), autoload_with=connection) + assert isinstance(reflected.c.numbers.type, types.ARRAY) + assert isinstance(reflected.c.numbers.type.item_type, types.Integer) + assert isinstance(reflected.c.labels.type.item_type, types.ARRAY) + assert reflected.c.amounts.type.item_type.precision == 30 + assert reflected.c.amounts.type.item_type.scale == 20 + for source in (table, reflected): + rows = connection.execute(select(source).order_by(source.c.id)).mappings().all() + eq_([dict(row) for row in rows], values) + result = connection.execute(select(table.c.labels).where(table.c.id == 1)) + cursor = getattr(result.context.cursor, "_cursor", result.context.cursor) + assert cursor.effective_engine_version == "Athena engine version 3" + + @sa_testing.combinations(False, True, argnames="literal_execute") + def test_typed_scalar_and_complex_elements(self, connection, literal_execute): + cases = [ + (AthenaArray(String), ["100%", "%(param_1)s", "back\\slash", "line\nbreak", "null"]), + (AthenaArray(types.Boolean), [True, False, None]), + (AthenaArray(types.Date), [date(2025, 1, 2), None]), + (AthenaArray(types.DateTime), [_datetime(2025, 1, 2, 3, 4, 5, 123000)]), + (AthenaArray(types.BINARY), [b"\x00\xff", b"", None]), + (AthenaArray(AthenaMap(Integer, String)), [{1: "001", 2: "a,b"}, {}, None]), + ( + AthenaArray(AthenaStruct(("name", String), ("n", Integer))), + [{"name": "a,b", "n": 2}, None], + ), + (AthenaArray(Integer, as_tuple=True), (1, None, 3)), + (AthenaArray(types.JSON), [{"n": 1, "s": "a,b", "fraction": 0.1}, [1, None], None]), + ] + expressions = [ + literal(value, type_=type_, literal_execute=literal_execute).label(f"v{index}") + for index, (type_, value) in enumerate(cases) + ] + row = connection.execute(select(*expressions)).one() + eq_(tuple(row), tuple(value for _, value in cases)) + + def test_nested_complex_reflection(self, connection, metadata): + table = Table( + "native_array_complex", + metadata, + Column( + "value", + AthenaArray( + AthenaStruct( + ("label", String), + ("numbers", AthenaArray(Integer)), + ("amount", types.Numeric(12, 3)), + ) + ), + ), + ) + table.create(connection) + value = [{"label": "001,a", "numbers": [1, None], "amount": Decimal("12.340")}] + connection.execute(table.insert().values(value=value)) + reflected = Table(table.name, MetaData(), autoload_with=connection) + item = reflected.c.value.type.item_type + assert isinstance(item, AthenaStruct) + assert isinstance(item.fields["numbers"], types.ARRAY) + assert item.fields["amount"].scale == 3 + eq_(connection.execute(select(reflected.c.value)).scalar_one(), value) + + class SimpleUpdateDeleteTest(_SimpleUpdateDeleteTest): @testing.variation("criteria", ["rows", "norows", "aggregate"]) @testing.requires.update_where_target_in_subquery From 5e86f013f150ee4357676d18e22fbd593ea7d6d2 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:35:46 +0900 Subject: [PATCH 06/20] Backport #775: Support SQLAlchemy ARRAY expressions in SELECT and WHERE (cherry picked from commit e083e8caa759a232f8a1a0b60d1cf086080ee841) Conflict resolution: The cherry-pick applied cleanly, but its tests in tests/sqlalchemy/test_suite.py use the 'sa_exc' alias for sqlalchemy.exc, which master imported through #771 (not backported). Added that import; the backported code is otherwise unchanged. Co-Authored-By: Claude Opus 5.5 --- docs/sqlalchemy.md | 34 +++++ pyathena/sqlalchemy/array.py | 54 ++++++- pyathena/sqlalchemy/compiler.py | 147 +++++++++++++++++- tests/pyathena/sqlalchemy/test_array.py | 85 +++++++++++ tests/pyathena/sqlalchemy/test_compiler.py | 144 ++++++++++++++++- tests/sqlalchemy/test_suite.py | 170 +++++++++++++++++++++ 6 files changed, 623 insertions(+), 11 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 9149bfa2..91aa6afe 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -1033,6 +1033,40 @@ CREATE TABLE orders ( ) ``` +#### ARRAY expressions + +Use SQLAlchemy expressions to index, slice, concatenate, or compare array values in SELECT and WHERE clauses. +Indices are one-based by default. +`AthenaArray(Integer, zero_indexes=True)` translates explicit indices and slice boundaries by adding one. +Reads return NULL when the resulting SQL index is NULL, below one, or out of range; negative indices do not count from the end. + +```python +from sqlalchemy import any_, select + +item_ids = orders.c.item_ids +statement = select(item_ids[1], item_ids[2:4], item_ids.concat([5])).where( + any_(item_ids) == 3 +) +``` + +Slice stops are inclusive, following SQLAlchemy's SQL array convention. +Omitted boundaries mean the beginning or end of the array; explicit boundaries are clipped to the array. +A reversed slice returns an empty array, and slicing a NULL array returns NULL. +Only `step=None` and `step=1` are supported. +Slice SQL can evaluate the array and boundary expressions more than once; use deterministic expressions rather than volatile functions such as `random()` in slices. +Use `AthenaArray` for open-ended slices with `zero_indexes=True`; SQLAlchemy's generic `ARRAY` comparator requires explicit bounds in that mode. +Indexing a multidimensional array retains the remaining dimensions, while slicing retains its array type. + +`any_(array)` and `all_(array)`, and the legacy `array.any(value)` and `array.all(value)` methods, compile to Athena's `any_match` and `all_match` functions. +Comparisons use SQL three-valued logic: NULL elements can produce NULL when no decisive true or false result exists. +For an empty array, ANY is false and ALL is true. +A NULL array produces NULL for both. +SQLAlchemy comparison flipping is preserved. +SQLAlchemy can rewrite negation into an element-wise comparison before dialect compilation, depending on the version and operand types. +For example, `~(any_(flags) == True)` can become an element-wise `!= True` comparison; older SQLAlchemy 2.0 releases also rewrite non-boolean comparisons this way. +To negate the whole match, always explicitly group the comparison first: `~(any_(flags) == True).self_group()`. +Quantifiers over subqueries retain their usual SQL compilation. + #### Querying ARRAY data Use `select()` with `ARRAY` or `AthenaArray` columns whose element types are known, either declared explicitly or reflected from Athena, to receive typed Python collections. diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index 17a19a90..cae77efb 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -1,4 +1,4 @@ -"""Athena ARRAY types, JSON result projection, and nested value processing.""" +"""Athena ARRAY types, expressions, JSON projection, and nested value processing.""" from __future__ import annotations @@ -6,11 +6,11 @@ from collections.abc import Mapping from datetime import date, datetime from decimal import Decimal -from typing import Any +from typing import TYPE_CHECKING, Any from sqlalchemy import cast, exc, types -from sqlalchemy.sql import sqltypes -from sqlalchemy.sql.elements import ColumnElement +from sqlalchemy.sql import operators, sqltypes +from sqlalchemy.sql.elements import ColumnElement, Slice from sqlalchemy.sql.type_api import TypeEngine from sqlalchemy.sql.visitors import InternalTraversal @@ -19,6 +19,12 @@ from pyathena.sqlalchemy.struct import AthenaStruct from pyathena.sqlalchemy.temporal import AthenaDate, AthenaTimestamp +# SQLAlchemy 2.0.0's ARRAY comparator is not generic at runtime. +if TYPE_CHECKING: + _ArrayComparatorBase = sqltypes.ARRAY.Comparator[Any] +else: + _ArrayComparatorBase = sqltypes.ARRAY.Comparator + class AthenaArray(sqltypes.ARRAY[Any]): """SQLAlchemy type for Athena ARRAY complex type. @@ -46,6 +52,24 @@ class AthenaArray(sqltypes.ARRAY[Any]): __visit_name__ = "array" + class Comparator(_ArrayComparatorBase): + """Build array indexing expressions with inclusive SQL slice bounds.""" + + def _setup_getitem(self, index): + if isinstance(index, slice): + if index.step is not None and (type(index.step) is not int or index.step != 1): + raise exc.CompileError("Athena ARRAY slices support only step=None or step=1") + start, stop = index.start, index.stop + if self.type.zero_indexes: + start = start + 1 if start is not None else None + stop = stop + 1 if stop is not None else None + return operators.getitem, Slice(start, stop, None), self.type + if self.type.zero_indexes: + index = index + 1 + return operators.getitem, index, _ArrayTypeInspector.item_type(self.type) + + comparator_factory = Comparator + def __init__( self, item_type: Any = None, @@ -93,6 +117,21 @@ def result_processor(self, dialect, coltype): return _ArrayValueProcessor(self, dialect).result +class _ArraySliceStepType(types.TypeDecorator[int]): + """Validate step values when SQLAlchemy reuses a generic ARRAY slice statement.""" + + impl = types.Integer + cache_ok = True + + def process_bind_param(self, value, dialect): + if type(value) is not int or value != 1: + raise ValueError("Athena ARRAY slices support only step=None or step=1") + return value + + def process_literal_param(self, value, dialect): + return self.process_bind_param(value, dialect) + + class ARRAY(AthenaArray): """Uppercase alias for AthenaArray type.""" @@ -137,6 +176,13 @@ class _ArrayTypeInspector: def __init__(self, dialect: Any) -> None: self.dialect = dialect + def array_type(self, type_: TypeEngine[Any]) -> sqltypes.ARRAY[Any] | None: + """Resolve the dialect's ARRAY implementation through variants and decorators.""" + implementation = type_.dialect_impl(self.dialect) + while isinstance(implementation, types.TypeDecorator): + implementation = self.decorator_impl(implementation) + return implementation if isinstance(implementation, sqltypes.ARRAY) else None + @staticmethod def item_type(type_: sqltypes.ARRAY[Any]) -> TypeEngine[Any]: if type_.dimensions is not None and type_.dimensions > 1: diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index aa0bbc80..feb54129 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -2,6 +2,7 @@ import re from collections.abc import Mapping +from itertools import product from typing import TYPE_CHECKING, Any, cast from sqlalchemy import exc, select, types, util @@ -16,20 +17,23 @@ from sqlalchemy.sql.elements import ( BindParameter, Cast, + CollectionAggregate, + Null, + Slice, TextClause, UnaryExpression, _label_reference, _textual_label_reference, ) from sqlalchemy.sql.schema import Column -from sqlalchemy.sql.selectable import CompoundSelect +from sqlalchemy.sql.selectable import CompoundSelect, ScalarSelect from pyathena.model import ( AthenaFileFormat, AthenaPartitionTransform, AthenaRowFormatSerde, ) -from pyathena.sqlalchemy.array import _ArrayTypeInspector +from pyathena.sqlalchemy.array import _ArraySliceStepType, _ArrayTypeInspector from pyathena.sqlalchemy.preparer import AthenaDDLIdentifierPreparer from pyathena.sqlalchemy.types import ( AthenaMap, @@ -259,6 +263,139 @@ def _array_type_inspector(self): def visit_char_length_func(self, fn: Function[Any], **kw: Any) -> str: return f"length{self.function_argspec(fn, **kw)}" + @staticmethod + def _original_froms(elements): + for element in elements: + while element._is_clone_of is not None: + element = element._is_clone_of + yield element + + def _array_lambda_name(self): + names = { + str( + element.text if isinstance(element, TextClause) else getattr(element, "name", "") + ).lower() + for element in visitors.iterate(self.statement) + } + index = getattr(self, "_array_lambda_index", 0) + # Textual SQL can embed a column name inside a larger expression. + while any(f"_pyathena_element_{index}" in name for name in names): + index += 1 + self._array_lambda_index = index + 1 + return f"_pyathena_element_{index}" + + def visit_binary( + self, + binary, + override_operator=None, + eager_grouping=False, + from_linter=None, + lateral_from_linter=None, + **kw, + ): + kw.update( + eager_grouping=eager_grouping, + from_linter=from_linter, + lateral_from_linter=lateral_from_linter, + ) + aggregate = binary.right + aggregate_on_left = isinstance(binary.left, CollectionAggregate) + if aggregate_on_left: + aggregate = binary.left + array_type = ( + self._array_type_inspector.array_type(aggregate.element.type) + if isinstance(aggregate, CollectionAggregate) + and not isinstance(aggregate.element, ScalarSelect) + else None + ) + if array_type is not None: + variable = self._array_lambda_name() + predicate = binary._clone() + item_type = _ArrayTypeInspector.item_type(array_type) + if aggregate_on_left: + predicate.left = Column(variable, item_type) + else: + predicate.right = Column(variable, item_type) + if isinstance(predicate.left, BindParameter) and ( + isinstance(item_type, types.ARRAY) + or ( + predicate.left.type is aggregate.element.type + and predicate.left.type._type_affinity is not types.ARRAY + ) + ): + predicate.left = predicate.left._with_binary_element_type(item_type) + if from_linter is not None and operators.is_comparison(binary.operator): + if lateral_from_linter is not None: + enclosing = [kw["enclosing_lateral"]] + lateral_from_linter.edges.update( + product( + self._original_froms(binary.left._from_objects + enclosing), + self._original_froms(binary.right._from_objects + enclosing), + ) + ) + else: + from_linter.edges.update( + product( + self._original_froms(binary.left._from_objects), + self._original_froms(binary.right._from_objects), + ) + ) + sql = super().visit_binary(predicate, override_operator=override_operator, **kw) + function = "any_match" if aggregate.operator is operators.any_op else "all_match" + array = self.process(aggregate.element, **kw) + return f"{function}({array}, {variable} -> {sql})" + return super().visit_binary(binary, override_operator=override_operator, **kw) + + def visit_getitem_binary(self, binary, operator, **kw): + array_type = self._array_type_inspector.array_type(binary.left.type) + if array_type is None: + raise exc.CompileError("Athena indexing requires an ARRAY expression") + array = self.process(binary.left, **kw) + if isinstance(binary.right, Slice): + bounds = binary.right + if not isinstance(bounds.step, Null) and not ( + isinstance(bounds.step, BindParameter) + and bounds.step.unique + and type(bounds.step.value) is int + and bounds.step.value == 1 + ): + raise exc.CompileError("Athena ARRAY slices support only step=None or step=1") + start = "1" if isinstance(bounds.start, Null) else self.process(bounds.start, **kw) + stop = ( + f"cardinality({array})" + if isinstance(bounds.stop, Null) + else self.process(bounds.stop, **kw) + ) + start = f"greatest({start}, 1)" + length = f"greatest(least({stop}, cardinality({array})) - {start} + 1, 0)" + sql = f"slice({array}, {start}, {length})" + return self._array_slice_step(sql, bounds.step, array_type, **kw) + index_expression = binary.right + if ( + isinstance(index_expression, BindParameter) + and self._array_type_inspector.array_type(index_expression.type) is not None + ): + index_expression = index_expression._with_binary_element_type(types.Integer()) + index = self.process(index_expression, **kw) + return f"element_at({array}, NULLIF(greatest({index}, 0), 0))" + + def _array_slice_step(self, sql, step, array_type, **kw): + if isinstance(step, Null): + return sql + if isinstance(step, BindParameter): + step = step._with_binary_element_type(_ArraySliceStepType()) + step_sql = self.process(step, **kw) + failure = ( + "CAST(concat('Unsupported ARRAY slice step: ', " + f"coalesce(CAST({step_sql} AS VARCHAR), 'NULL')) AS BIGINT)" + ) + empty = ( + f"slice({sql}, 1, 0)" + if _ArrayTypeInspector.has_unknown_element(array_type) + else f"CAST(ARRAY[] AS {self._complex_dml_type(array_type)})" + ) + return f"IF({step_sql} = 1, {sql}, slice({empty}, {failure}, 0))" + def translate_select_structure(self, select_stmt, **kw): """Keep DISTINCT and ordering on native arrays before result serialization.""" if ( @@ -287,10 +424,8 @@ def visit_compound_select(self, cs, asfrom=False, compound_index=None, **kw): return super().visit_compound_select(cs, asfrom=asfrom, compound_index=compound_index, **kw) def _has_array_result(self, column): - type_ = column.type.dialect_impl(self.dialect) - while isinstance(type_, types.TypeDecorator): - type_ = self._array_type_inspector.decorator_impl(type_) - return isinstance(type_, types.ARRAY) and not _ArrayTypeInspector.has_unknown_element(type_) + type_ = self._array_type_inspector.array_type(column.type) + return type_ is not None and not _ArrayTypeInspector.has_unknown_element(type_) def _array_result_select(self, statement): if any( diff --git a/tests/pyathena/sqlalchemy/test_array.py b/tests/pyathena/sqlalchemy/test_array.py index 5dce0576..e52498c1 100644 --- a/tests/pyathena/sqlalchemy/test_array.py +++ b/tests/pyathena/sqlalchemy/test_array.py @@ -13,7 +13,11 @@ MetaData, String, Table, + all_, + any_, + bindparam, cast, + column, literal, literal_column, select, @@ -208,7 +212,88 @@ def test_array_reflection_warns_for_unrecognized_nested_type(self, signature): assert isinstance(type_, types.NullType) +class TestAthenaArrayComparator: + @staticmethod + def _compile_sql(expression): + return str( + expression.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True}) + ) + + @pytest.mark.parametrize("array_type", [AthenaArray(Integer), types.ARRAY(Integer)]) + @pytest.mark.parametrize("index", [-2, 0, 1, 100]) + def test_array_index(self, array_type, index): + value = column("items", array_type) + assert ( + self._compile_sql(value[index]) == f"element_at(items, NULLIF(greatest({index}, 0), 0))" + ) + assert isinstance(value[index].type, Integer) + + def test_array_dimensions_and_zero_indexes(self): + value = column( + "items", AthenaArray(Integer, dimensions=3, zero_indexes=True, as_tuple=True) + ) + assert value[0].type.dimensions == 2 + assert value[0][0].type.dimensions == 1 + assert isinstance(value[0][0][0].type, Integer) + assert value[0].type.as_tuple + assert value[0].type.zero_indexes + assert "NULLIF(greatest(1, 0), 0)" in self._compile_sql(value[0]) + assert "greatest(1, 1)" in self._compile_sql(value[:0]) + assert "least(1, cardinality(items))" in self._compile_sql(value[:0]) + assert "cardinality(items)" in self._compile_sql(value[0:]) + nested = column("nested", AthenaArray(AthenaArray(String))) + assert isinstance(nested[1].type, AthenaArray) + assert isinstance(nested[1][1].type, String) + + @pytest.mark.parametrize( + "bounds", [slice(None), slice(1, 2), slice(-2, 100), slice(3, 1), slice(1, 3, 1)] + ) + def test_array_slice(self, bounds): + value = column("items", AthenaArray(Integer)) + result = value[bounds] + assert result.type is value.type + sql = self._compile_sql(result) + assert sql.startswith("slice(items, greatest(") + assert "greatest(least(" in sql + + @pytest.mark.parametrize("step", [0, 2, -1, True, 1.0, bindparam("step", 1)]) + def test_array_slice_rejects_steps(self, step): + with pytest.raises(sa_exc.CompileError, match="step"): + self._compile_sql(column("items", AthenaArray(Integer))[1:3:step]) + + def test_decorated_array_index_and_slice(self): + items = column("items", TupleArray()) + assert self._compile_sql(items[1]) == "element_at(items, NULLIF(greatest(1, 0), 0))" + assert isinstance(items[1].type, Integer) + compiled = select(items[1:2]).compile(dialect=AthenaDialect()) + assert "transform(slice(items," in str(compiled) + processor = ( + compiled._result_columns[0] + .type.dialect_impl(AthenaDialect()) + .result_processor(AthenaDialect(), None) + ) + assert processor('{"_pyathena_array":["1","2"]}') == (1, 2) + + class TestArrayTypeInspector: + @pytest.mark.parametrize( + ("type_", "value", "expected"), + [ + (TupleArray(), 2, "2"), + (String().with_variant(TupleArray(), "awsathena"), 2, "2"), + (Integer().with_variant(TupleArray(), "awsathena"), 2, "2"), + (String().with_variant(AthenaArray(String), "awsathena"), "a", "'a'"), + ], + ) + def test_decorated_and_variant_array_quantifiers(self, type_, value, expected): + items = column("items", type_) + statement = select(any_(items) == value, all_(items) > value) + sql = str( + statement.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True}) + ) + assert f"any_match((items), _pyathena_element_0 -> {expected} = _pyathena_element_0)" in sql + assert f"all_match((items), _pyathena_element_1 -> {expected} < _pyathena_element_1)" in sql + @pytest.mark.parametrize( ("type_", "ddl", "dml"), [ diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index 16c9bbba..a160cfc1 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -1,3 +1,5 @@ +import warnings + import pytest from sqlalchemy import ( Column, @@ -8,12 +10,21 @@ Numeric, String, Table, + all_, + any_, + bindparam, + cast, + column, exc, func, select, + table, + text, + types, ) from sqlalchemy.engine.url import make_url -from sqlalchemy.sql import literal, literal_column +from sqlalchemy.sql import literal, literal_column, operators +from sqlalchemy.sql.compiler import FROM_LINTING from sqlalchemy.sql.ddl import CreateTable from pyathena.sqlalchemy.base import AthenaDialect @@ -272,6 +283,137 @@ def test_visit_truediv_binary(self, expression, expected): assert str(compiled) == expected + def _compile_sql(self, expression): + return str(expression.compile(dialect=self.dialect, compile_kwargs={"literal_binds": True})) + + @pytest.mark.parametrize(("aggregate", "function"), [(any_, "any_match"), (all_, "all_match")]) + @pytest.mark.parametrize( + ("op", "sql_operator"), + [ + (operators.eq, "="), + (operators.ne, "!="), + (operators.lt, "<"), + (operators.le, "<="), + (operators.gt, ">"), + (operators.ge, ">="), + ], + ) + def test_array_quantified_comparison(self, aggregate, function, op, sql_operator): + items = column("items", AthenaArray(Integer)) + sql = self._compile_sql(op(2, aggregate(items))) + assert sql == ( + f"{function}((items), _pyathena_element_0 -> 2 {sql_operator} _pyathena_element_0)" + ) + + def test_array_quantifier_null_negation_and_legacy_methods(self): + items = column("items", AthenaArray(Integer)) + assert "NULL = _pyathena_element_0" in self._compile_sql(any_(items) == None) # noqa: E711 + assert self._compile_sql(items.any(2)) == self._compile_sql(any_(items) == 2) + assert self._compile_sql(items.all(2, operator=operators.lt)) == self._compile_sql( + all_(items) > 2 + ) + assert self._compile_sql(~items.any(2)).startswith("NOT (any_match(") + assert "2 > _pyathena_element_0" in self._compile_sql(any_(items) < 2) + + @pytest.mark.parametrize( + ("scalar", "expected"), + [ + (column("_pyathena_element_0", Integer), "_pyathena_element_0"), + (literal_column("_pyathena_element_0 + 1"), "_pyathena_element_0 + 1"), + (literal_column('"_pyathena_element_0" + 1'), '"_pyathena_element_0" + 1'), + (text("_pyathena_element_0 + 1"), "_pyathena_element_0 + 1"), + ], + ) + def test_array_lambda_does_not_capture_column_names(self, scalar, expected): + items = column("items", AthenaArray(Integer)) + sql = self._compile_sql(scalar == any_(items)) + assert f"_pyathena_element_1 -> {expected} = _pyathena_element_1" in sql + + def test_array_index_does_not_repeat_expression(self): + items = column("items", AthenaArray(Integer)) + index = cast(func.floor(func.random() * 3), Integer) - 1 + sql = self._compile_sql(items[index]) + assert sql.count("random()") == 1 + assert "NULLIF(greatest(" in sql + + def test_multidimensional_array_quantifier_bind_type(self): + items = column("items", AthenaArray(Integer, dimensions=2)) + assert "CAST(ARRAY[1, 2] AS ARRAY(INTEGER)) =" in self._compile_sql(items.any([1, 2])) + + def test_quantifier_preserves_explicit_array_bind(self): + items = column("items", AthenaArray(Integer)) + needle = literal([1, 2], items.type) + sql = self._compile_sql(needle == any_(func.array_agg(items))) + assert "CAST(ARRAY[1, 2] AS ARRAY(INTEGER)) = _pyathena_element_0" in sql + unknown = column("unknown", AthenaArray(types.NullType())) + sql = self._compile_sql(needle == any_(unknown)) + assert "CAST(ARRAY[1, 2] AS ARRAY(INTEGER)) = _pyathena_element_0" in sql + + def test_subquery_any_remains_unchanged(self): + sql = self._compile_sql(any_(select(column("item", Integer)).scalar_subquery()) == 2) + assert "ANY (SELECT item)" in sql + assert "any_match" not in sql + items = column("items", AthenaArray(Integer)) + array_subquery = self._compile_sql(select(items == any_(select(items).scalar_subquery()))) + assert "ANY (SELECT items" in array_subquery + assert "any_match" not in array_subquery + + def test_array_concat_and_cache_bind_values(self): + items = column("items", AthenaArray(Integer)) + assert self._compile_sql(items.concat([2])) == "items || CAST(ARRAY[2] AS ARRAY(INTEGER))" + first = select(items[bindparam("index")]).where(any_(items) == 2) + second = select(items[bindparam("index")]).where(any_(items) == 3) + assert first._generate_cache_key().key == second._generate_cache_key().key + assert self._compile_sql(first.params(index=1)) != self._compile_sql(first.params(index=2)) + + def test_quantifier_boolean_left_operand(self): + flags = column("flags", AthenaArray(types.Boolean)) + assert "any_match" in self._compile_sql(any_(flags) == True) # noqa: E712 + assert "all_match" in self._compile_sql(all_(flags) != False) # noqa: E712 + assert "IS DISTINCT FROM" in self._compile_sql(any_(flags).is_distinct_from(True)) + comparison = any_(flags) == True # noqa: E712 + assert self._compile_sql(~comparison) == ( + "any_match((flags), _pyathena_element_0 -> _pyathena_element_0 != true)" + ) + assert self._compile_sql(~comparison.self_group()) == ( + "NOT (any_match((flags), _pyathena_element_0 -> _pyathena_element_0 = true))" + ) + + def test_quantifier_join_linter_tracks_original_tables(self): + left = table("left_table", column("value", Integer)) + right = table("right_table", column("items", AthenaArray(Integer))) + query = select(left, right).where(left.c.value == any_(right.c["items"])) + with warnings.catch_warnings(): + warnings.simplefilter("error", exc.SAWarning) + query.compile(dialect=AthenaDialect(), linting=FROM_LINTING) + + def test_generic_array_slice_step_is_rendered_for_cache_validation(self): + items = column("items", types.ARRAY(Integer)) + query = select(items[2:3:1]) + compiled = query.compile(dialect=AthenaDialect()) + assert "Unsupported ARRAY slice step" in str(compiled) + assert list(compiled.params.values()).count(1) == 1 + assert query._generate_cache_key().key == select(items[2:3:2])._generate_cache_key().key + step_name = next(name for name, value in compiled.params.items() if value == 1) + assert f"IF(%({step_name})s = 1" in str(compiled) + for invalid in (2, 1.0, None): + with pytest.raises(ValueError, match="step"): + compiled._bind_processors[step_name](invalid) + native = column("items", AthenaArray(Integer)) + assert ( + select(native[1:3:1])._generate_cache_key().key + == select(native[1:3])._generate_cache_key().key + ) + with pytest.raises(exc.CompileError, match="step"): + native[1:3:2] + + def test_stepped_array_aggregate_with_inferred_element_type(self): + expression = func.array_agg(func.length("abc"))[1:2:1] + sql = self._compile_sql(select(expression)) + assert "array_agg(length('abc'))" in sql + assert "Unsupported ARRAY slice step" in sql + assert "ARRAY(NULL)" not in sql + class TestAthenaDDLCompiler: """Compile-only (no AWS) tests for the DDL compiler's S3 Tables support. diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index f9c9d89b..b19fe7a9 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -8,6 +8,9 @@ Integer, MetaData, String, + all_, + any_, + bindparam, cast, func, literal, @@ -18,6 +21,7 @@ types, ) from sqlalchemy import Table as SATable +from sqlalchemy import exc as sa_exc from sqlalchemy import testing as sa_testing from sqlalchemy.testing import eq_, fixtures from sqlalchemy.testing.schema import Column, Table @@ -122,6 +126,172 @@ def process_result_value(self, value, dialect): return tuple(value) if value is not None else None +class ArrayExpressionTest(fixtures.TestBase): + __backend__ = True + __requires__ = ("array_type",) + + def test_decorated_and_variant_arrays(self, connection): + array = literal([1, 2, 3], _ArrayTuple()) + variant = literal([1, 2], String().with_variant(_ArrayTuple(), "awsathena")) + statement = select( + array[1], + array[1:2], + any_(array) == 2, + all_(array) > 0, + array.concat([4]), + any_(variant) == 2, + variant, + ) + eq_( + tuple(connection.execute(statement).one()), + (1, (1, 2), True, True, (1, 2, 3, 4), True, (1, 2)), + ) + + @sa_testing.combinations( + (String(), String, ["a", "b"], "a"), + (Integer(), Integer, [1, 2], 1), + argnames="base_type,item_type,values,needle", + ) + def test_variant_array_quantifier_scalar_bind( + self, connection, base_type, item_type, values, needle + ): + array = literal(values, base_type.with_variant(AthenaArray(item_type), "awsathena")) + statement = select(any_(array) == needle, all_(array) == needle) + eq_(tuple(connection.execute(statement).one()), (True, False)) + eq_(tuple(connection.execute(statement).one()), (True, False)) + + def test_cached_steps_and_boolean_quantifiers(self, connection): + inferred = func.array_agg(func.length(literal("abc"))) + eq_(connection.execute(select(inferred[1:2:1])).scalar_one(), [3]) + array = literal([1, 2, 3], types.ARRAY(Integer)) + eq_(connection.execute(select(array[1:2:1])).scalar_one(), [1, 2]) + with pytest.raises(sa_exc.StatementError, match="step") as error: + connection.execute(select(array[1:2:2])).all() + assert isinstance(error.value.orig, ValueError) + flags = literal([True, False], AthenaArray(types.Boolean)) + comparison = any_(flags) == True # noqa: E712 + eq_( + tuple( + connection.execute( + select(comparison, all_(flags) == True, ~comparison, ~comparison.self_group()) # noqa: E712 + ).one() + ), + (True, False, True, False), + ) + + def test_index_slice_and_concat(self, connection): + array = literal([1, None, 3], AthenaArray(Integer)) + empty = literal([], AthenaArray(Integer)) + missing = literal(None, AthenaArray(Integer)) + expressions = [ + array[1], + array[2], + array[0], + array[-1], + array[10], + array[literal(None, Integer)], + array[:], + array[:2], + array[2:], + array[-2:100], + array[3:1], + empty[1:3], + missing[1:3], + array.concat([4]), + array[1:2].concat(array[3:]), + ] + eq_( + tuple(connection.execute(select(*expressions)).one()), + ( + 1, + None, + None, + None, + None, + None, + [1, None, 3], + [1, None], + [None, 3], + [1, None, 3], + [], + [], + None, + [1, None, 3, 4], + [1, None, 3], + ), + ) + + def test_dimensions_zero_indexes_and_bound_index(self, connection): + array = literal([[1, 2], [3]], AthenaArray(Integer, dimensions=2, zero_indexes=True)) + statement = select(array[bindparam("index")], array[:0], array[0][1], array[1:]) + eq_(tuple(connection.execute(statement, {"index": 0}).one()), ([1, 2], [[1, 2]], 2, [[3]])) + eq_(tuple(connection.execute(statement, {"index": 1}).one()), ([3], [[1, 2]], 2, [[3]])) + eq_(connection.execute(select(array.any([1, 2]))).scalar_one(), True) + plain = literal([1, 2], types.ARRAY(Integer)) + eq_(connection.execute(select(plain[bindparam("index")]), {"index": 2}).scalar_one(), 2) + eq_( + connection.execute( + select(literal([1, 2], AthenaArray(Integer)) == any_(func.array_agg(plain))) + ).scalar_one(), + True, + ) + + def test_quantified_comparisons(self, connection): + cases = [ + ([1, 2, None], True, False, False, False), + ([1, None], None, False, False, None), + ([2, None], True, None, False, False), + ([3, None], None, False, None, None), + ([1, 3], False, False, False, True), + ([], False, True, True, True), + (None, None, None, None, None), + ] + for values, eq_any, eq_all, lt_all, neg_any in cases: + array = literal(values, AthenaArray(Integer)) + row = connection.execute( + select( + any_(array) == 2, + all_(array) == 2, + all_(array) > 2, + ~array.any(2), + any_(array) == None, # noqa: E711 + ) + ).one() + eq_( + tuple(row), + ( + eq_any, + eq_all, + lt_all, + neg_any if eq_any is not None else None, + False if values == [] else None, + ), + ) + + def test_where_and_lambda_names(self, connection, metadata): + table = Table( + "array_expressions", + metadata, + Column("id", Integer), + Column("items", AthenaArray(Integer)), + Column("_pyathena_element_0", Integer), + ) + table.create(connection) + connection.execute( + table.insert(), + [ + {"id": 1, "items": [1, 3], "_pyathena_element_0": 3}, + {"id": 2, "items": [2, 4], "_pyathena_element_0": 1}, + ], + ) + predicate = (table.c._pyathena_element_0 == any_(table.c["items"])) & ( + table.c["items"][1] == 1 + ) + eq_(connection.execute(select(table.c.id).where(predicate)).scalars().all(), [1]) + textual_predicate = literal_column("_pyathena_element_0 + 1") == any_(table.c["items"]) + eq_(connection.execute(select(table.c.id).where(textual_predicate)).scalars().all(), [2]) + + class NativeArrayTest(fixtures.TestBase): __backend__ = True __requires__ = ("array_type",) From e834e924ae2f3de22f7790f4ef1eec4f3686f36c Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:36:04 +0900 Subject: [PATCH 07/20] Backport #776: Support SQLAlchemy ARRAY partial UPDATE and slice resizing (cherry picked from commit 15325e172714e2e809dd5c597e8babb2c51a6fac) Conflict resolution: tests/sqlalchemy/test_suite.py conflicted because #771 and #777 (not backported) changed its imports. Applied only #776's own changes: import update, and Session and registry from sqlalchemy.orm, and add ArrayUpdateTest verbatim in master's position between _ArrayTuple and ArrayExpressionTest. Co-Authored-By: Claude Opus 5.5 --- docs/sqlalchemy.md | 41 ++++ pyathena/formatter.py | 2 +- pyathena/sqlalchemy/array.py | 264 +++++++++++++++++++++++- pyathena/sqlalchemy/compiler.py | 38 +++- tests/pyathena/sqlalchemy/test_array.py | 237 +++++++++++++++++++++ tests/sqlalchemy/test_suite.py | 213 +++++++++++++++++++ 6 files changed, 775 insertions(+), 20 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 91aa6afe..e3229472 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -1067,6 +1067,47 @@ For example, `~(any_(flags) == True)` can become an element-wise `!= True` compa To negate the whole match, always explicitly group the comparison first: `~(any_(flags) == True).self_group()`. Quantifiers over subqueries retain their usual SQL compilation. +#### Partial ARRAY updates + +Use indexed or sliced columns as UPDATE assignment keys on Iceberg tables. +PyAthena compiles each assignment into a single server-side UPDATE of the whole array. +Other columns can be assigned in the same statement. + +For an existing Iceberg table named `orders` with an ARRAY column `item_ids`: + +```python +from sqlalchemy import MetaData, Table + +orders = Table("orders", MetaData(), autoload_with=engine) +item_ids = orders.c.item_ids +with engine.begin() as conn: + conn.execute(orders.update().values({item_ids[2]: 10})) + conn.execute(orders.update().values({item_ids[2:3]: [20, 30, 40]})) +``` + +Element assignment beyond the end extends the array and fills intervening positions with NULL. +NULL padding uses Athena's `repeat()` function and follows its size limits. +A NULL destination array is treated as empty for partial updates. +Assigning `None` to an element stores NULL. +Nested indices rebuild the corresponding inner arrays, including missing inner arrays. +The same one-based default and `zero_indexes=True` translation apply to reads and writes. + +Slice assignment replaces an inclusive range with any number of elements. +An empty replacement array deletes the range; a longer or shorter replacement resizes the array. +A reversed range inserts before its start position. +A start beyond the end pads with NULL before inserting, and a stop beyond the end only removes existing elements. +Omitted boundaries mean the beginning or end. +These resize rules are PyAthena-specific and do not promise full PostgreSQL array-assignment compatibility. + +Write indices and explicit slice boundaries must be non-NULL positive integers after normalization. +A slice replacement must be a non-NULL array; use `[]` to delete elements. +Decimal partial updates require a declared `Numeric` precision, including SQL-expression assignments; specify a scale to retain fractional values. +Validation can raise `CompileError` during compilation, `StatementError` during binding, or `DBAPIError` from Athena; the exception class can differ when a compiled statement is reused. +Only `step=None` and `step=1` are supported, and only the final component of a nested update path may be a slice. +PyAthena rejects multiple partial assignments to the same array column, or a partial assignment combined with a whole-column assignment to that column. +Use one whole-array expression when an update needs several changes to the same array. +Partial updates can evaluate indices, boundaries, and replacement SQL expressions more than once; use deterministic expressions. + #### Querying ARRAY data Use `select()` with `ARRAY` or `AthenaArray` columns whose element types are known, either declared explicitly or reflected from Athena, to receive typed Python collections. diff --git a/pyathena/formatter.py b/pyathena/formatter.py index f4f5caea..0c6a7c13 100644 --- a/pyathena/formatter.py +++ b/pyathena/formatter.py @@ -22,7 +22,7 @@ class _ComplexParameter: """Typed complex value supplied by the SQLAlchemy dialect.""" - constructor: Literal["ARRAY", "MAP", "ROW", "JSON_PARSE"] + constructor: Literal["ARRAY", "MAP", "ROW", "JSON_PARSE", "FROM_HEX"] values: tuple[Any, ...] diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index cae77efb..09cefe69 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -8,9 +8,10 @@ from decimal import Decimal from typing import TYPE_CHECKING, Any -from sqlalchemy import cast, exc, types +from sqlalchemy import cast, exc, types, util from sqlalchemy.sql import operators, sqltypes -from sqlalchemy.sql.elements import ColumnElement, Slice +from sqlalchemy.sql.elements import BinaryExpression, BindParameter, ColumnElement, Null, Slice +from sqlalchemy.sql.schema import Column from sqlalchemy.sql.type_api import TypeEngine from sqlalchemy.sql.visitors import InternalTraversal @@ -220,23 +221,23 @@ def has_unknown_element(type_: TypeEngine[Any]) -> bool: class _ArrayValueProcessor: - """Convert one declared ARRAY type between Python values and Athena transport. + """Convert ARRAY values and their typed elements to and from Athena transport. SQLAlchemy constructs processors per type and dialect. Keep that context here and share the recursive ARRAY/MAP/ROW traversal across bind parameters, SQL literals, and fetched JSON results. """ - def __init__(self, array_type: AthenaArray, dialect: Any) -> None: - self.array_type = array_type + def __init__(self, type_: TypeEngine[Any], dialect: Any) -> None: + self.type_ = type_ self.dialect = dialect self._type_inspector = _ArrayTypeInspector(dialect) def bind(self, value: Any) -> Any: - return self._bind(value, self.array_type) + return self._bind(value, self.type_) def literal(self, value: Any) -> str: - return self._literal(value, self.array_type) + return self._literal(value, self.type_) def result(self, value: Any) -> Any: if value is None: @@ -250,7 +251,11 @@ def result(self, value: Any) -> Any: return value if isinstance(value, dict) and "_pyathena_array" in value: value = value["_pyathena_array"] - return self._decode(value, self.array_type, self.array_type.as_tuple) + return self._decode( + value, + self.type_, + self.type_.as_tuple if isinstance(self.type_, sqltypes.ARRAY) else False, + ) @staticmethod def _complex_values(value: Any, type_: TypeEngine[Any]): @@ -399,3 +404,246 @@ def _decode(self, value: Any, type_: TypeEngine[Any], as_tuple: bool = False) -> processor = type_.dialect_impl(self.dialect).result_processor(self.dialect, None) return processor(value) if processor else value return value + + +class _ArrayAssignmentType(types.TypeDecorator[Any]): + """Preserve declared element processors for ARRAY assignment values.""" + + impl = types.NullType + cache_ok = True + + def __init__(self, item_type): + super().__init__() + self.item_type = item_type + + def bind_processor(self, dialect): + processor = _ArrayValueProcessor(self.item_type, dialect) + + def process(value): + value = processor.bind(value) + if isinstance(value, (bytes, bytearray)): + return _ComplexParameter("FROM_HEX", (value.hex(),)) + return value + + return process + + def literal_processor(self, dialect): + return _ArrayValueProcessor(self.item_type, dialect).literal + + def bind_expression(self, bindvalue): + expression = self.item_type.bind_expression(bindvalue) + return bindvalue if expression is None else expression + + +class _ArrayWriteIndexType(types.TypeDecorator[int]): + """Reject non-integer and NULL bound ARRAY write indices.""" + + impl = types.Integer + cache_ok = True + + def process_bind_param(self, value, dialect): + if type(value) is not int: + raise ValueError("ARRAY write indices must be non-NULL integers") + return value + + def process_literal_param(self, value, dialect): + return self.process_bind_param(value, dialect) + + +class _ArrayUpdate(ColumnElement[Any]): + """Whole-column expression generated from one partial ARRAY assignment.""" + + __visit_name__ = "athena_array_update" + inherit_cache = True + _traverse_internals = [ # noqa: RUF012 + ("column", InternalTraversal.dp_clauseelement), + ("path", InternalTraversal.dp_clauseelement_list), + ("value", InternalTraversal.dp_clauseelement), + ("type", InternalTraversal.dp_type), + ] + + def __init__(self, column, path, value, value_type): + self.column = column + self.path = path + self.type = column.type + self.value_type = value_type + self.value = ( + value._with_binary_element_type( + _ArrayAssignmentType(value_type if value.type._isnull else value.type) + ) + if isinstance(value, BindParameter) + else value + ) + + @property + def _from_objects(self): + return self.column._from_objects + self.value._from_objects + + @classmethod + def rewrite(cls, statement, dialect): + inspector = _ArrayTypeInspector(dialect) + values = statement._ordered_values + if values is None: + values = list((statement._values or {}).items()) + rewritten = [] + seen = set() + partial = set() + for key, value in values: + base = key + path: list[Any] = [] + while isinstance(base, BinaryExpression) and base.operator is operators.getitem: + if inspector.array_type(base.left.type) is None: + break + path.insert(0, base.right) + base = base.left + name = base if isinstance(base, str) else getattr(base, "key", None) + if path: + if ( + not isinstance(base, Column) + or base.table is None + or base.table._deannotate() is not statement.table._deannotate() + ): + raise exc.CompileError("ARRAY updates require a column of the target table") + if name in seen: + raise exc.CompileError("Only one assignment per ARRAY column is supported") + if any(isinstance(index, Slice) for index in path[:-1]): + raise exc.CompileError("Only the final ARRAY update index can be a slice") + partial.add(name) + value_type = key.type + value = cls(base, path, value, value_type) + key = base + elif name in partial: + raise exc.CompileError("Only one assignment per ARRAY column is supported") + seen.add(name) + rewritten.append((key, value)) + if not partial: + return statement + result = statement._clone() + if statement._ordered_values is not None: + result._ordered_values = rewritten + else: + result._values = util.immutabledict(rewritten) + return result + + +class _ArrayUpdateCompiler: + """Render an ARRAY assignment by rebuilding its affected nested arrays.""" + + def __init__(self, compiler): + self.compiler = compiler + self._type_inspector = _ArrayTypeInspector(compiler.dialect) + + def process(self, expression, **kw): + compiler = self.compiler + value = expression.value + final_slice = isinstance(expression.path[-1], Slice) + if final_slice and ( + isinstance(value, Null) + or ( + isinstance(value, BindParameter) + and not value.required + and value.callable is None + and value.value is None + ) + ): + raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array") + rhs = compiler.process(value, **kw) + rhs_type = compiler._complex_dml_type(expression.value_type, require_precision=True) + rhs = f"CAST({rhs} AS {rhs_type})" + if final_slice: + # Reject SQL expressions that evaluate to NULL without issuing a second statement. + failure = ( + f"slice(CAST(ARRAY[] AS {rhs_type}), " + "CAST(concat('NULL ARRAY slice assignment', coalesce(CAST(cardinality(" + f"{rhs}) AS VARCHAR), '')) AS BIGINT), 0)" + ) + rhs = f"IF({rhs} IS NULL, {failure}, {rhs})" + return self._rebuild( + compiler.process(expression.column, **kw), expression.type, expression.path, rhs, **kw + ) + + def _index_sql(self, index: ColumnElement[Any], **kw): + compiler = self.compiler + if isinstance(index, Null): + raise exc.CompileError("ARRAY write indices must be non-NULL positive integers") + if ( + isinstance(index, BindParameter) + and not index.required + and index.callable is None + and (type(index.value) is not int or index.value <= 0) + ): + raise exc.CompileError( + "ARRAY write indices must be positive integers after normalization" + ) + if not isinstance(index.type, (types.Integer, types.NullType)) and not ( + isinstance(index, BindParameter) + and self._type_inspector.array_type(index.type) is not None + ): + raise exc.CompileError("ARRAY write indices must be integers") + + if isinstance(index, BindParameter): + index = index._with_binary_element_type(_ArrayWriteIndexType()) + sql = compiler.process(index, **kw) + failure = ( + "CAST(concat('Invalid ARRAY index: ', " + f"coalesce(CAST({sql} AS VARCHAR), 'NULL')) AS BIGINT)" + ) + return f"IF({sql} > 0, {sql}, {failure})" + + def _rebuild(self, array, array_type, path, rhs, **kw): + compiler = self.compiler + array_type = self._type_inspector.array_type(array_type) + if array_type is None: + raise exc.CompileError("Partial ARRAY updates require an ARRAY column type") + array_sql_type = compiler._complex_dml_type(array_type) + array = f"coalesce({array}, CAST(ARRAY[] AS {array_sql_type}))" + bound = path[0] + if isinstance(bound, Slice): + return self._rebuild_slice(array, array_type, bound, rhs, **kw) + return self._rebuild_element(array, array_type, bound, path[1:], rhs, **kw) + + def _prefix_and_padding(self, array, start, array_type): + prefix = f"slice({array}, 1, least({start} - 1, cardinality({array})))" + element_type = self.compiler._complex_dml_type(_ArrayTypeInspector.item_type(array_type)) + padding = ( + f"repeat(CAST(NULL AS {element_type}), " + f"CAST(greatest({start} - 1 - cardinality({array}), 0) AS INTEGER))" + ) + return prefix, padding + + def _rebuild_slice(self, array, array_type, bound, rhs, **kw): + if not isinstance(bound.step, Null) and not ( + isinstance(bound.step, BindParameter) + and bound.step.unique + and type(bound.step.value) is int + and bound.step.value == 1 + ): + raise exc.CompileError("Athena ARRAY slices support only step=None or step=1") + start = "1" if isinstance(bound.start, Null) else self._index_sql(bound.start, **kw) + stop = ( + f"cardinality({array})" + if isinstance(bound.stop, Null) + else self._index_sql(bound.stop, **kw) + ) + prefix, padding = self._prefix_and_padding(array, start, array_type) + tail_start = f"greatest({start}, {stop} + 1)" + suffix = ( + f"slice({array}, {tail_start}, greatest(cardinality({array}) - {tail_start} + 1, 0))" + ) + return self.compiler._array_slice_step( + f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, array_type, **kw + ) + + def _rebuild_element(self, array, array_type, bound, remaining_path, rhs, **kw): + index = self._index_sql(bound, **kw) + previous = f"element_at({array}, {index})" + replacement = ( + self._rebuild( + previous, _ArrayTypeInspector.item_type(array_type), remaining_path, rhs, **kw + ) + if remaining_path + else rhs + ) + prefix, padding = self._prefix_and_padding(array, index, array_type) + suffix = f"slice({array}, {index} + 1, greatest(cardinality({array}) - {index}, 0))" + return f"concat({prefix}, {padding}, ARRAY[{replacement}], {suffix})" diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index feb54129..3ca46bfa 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -33,7 +33,12 @@ AthenaPartitionTransform, AthenaRowFormatSerde, ) -from pyathena.sqlalchemy.array import _ArraySliceStepType, _ArrayTypeInspector +from pyathena.sqlalchemy.array import ( + _ArraySliceStepType, + _ArrayTypeInspector, + _ArrayUpdate, + _ArrayUpdateCompiler, +) from pyathena.sqlalchemy.preparer import AthenaDDLIdentifierPreparer from pyathena.sqlalchemy.types import ( AthenaMap, @@ -270,6 +275,15 @@ def _original_froms(elements): element = element._is_clone_of yield element + def visit_update(self, update_stmt, visiting_cte=None, **kw): + """Rewrite partial array assignments into one native Athena UPDATE.""" + return super().visit_update( + _ArrayUpdate.rewrite(update_stmt, self.dialect), visiting_cte=visiting_cte, **kw + ) + + def visit_athena_array_update(self, expression, **kw): + return _ArrayUpdateCompiler(self).process(expression, **kw) + def _array_lambda_name(self): names = { str( @@ -620,7 +634,7 @@ def visit_truediv_binary(self, binary, operator, **kw): def visit_cast(self, cast: Cast[Any], **kwargs): if isinstance(cast.type, (types.ARRAY, AthenaMap, AthenaStruct)): type_clause = self._complex_dml_type( - cast.type, implicit_bind=cast._annotations.get("_pyathena_array_bind", False) + cast.type, require_precision=cast._annotations.get("_pyathena_array_bind", False) ) return f"CAST({self.process(cast.clause, **kwargs)} AS {type_clause})" if (isinstance(cast.type, types.VARCHAR) and cast.type.length is None) or isinstance( @@ -642,27 +656,29 @@ def visit_cast(self, cast: Cast[Any], **kwargs): type_clause = cast.typeclause._compiler_dispatch(self, **kwargs) return f"CAST({cast.clause._compiler_dispatch(self, **kwargs)} AS {type_clause})" - def _complex_dml_type(self, type_, *, implicit_bind=False): + def _complex_dml_type(self, type_, *, require_precision=False): if isinstance(type_, types.TypeDecorator): return self._complex_dml_type( - self._array_type_inspector.decorator_impl(type_), implicit_bind=implicit_bind + self._array_type_inspector.decorator_impl(type_), + require_precision=require_precision, ) if isinstance(type_, types.NullType): raise exc.CompileError("Bound ARRAY values require an explicit element type") if isinstance(type_, types.ARRAY): item = self._complex_dml_type( - _ArrayTypeInspector.item_type(type_), implicit_bind=implicit_bind + _ArrayTypeInspector.item_type(type_), require_precision=require_precision ) return f"ARRAY({item})" if isinstance(type_, AthenaMap): - return ( - f"MAP({self._complex_dml_type(type_.key_type, implicit_bind=implicit_bind)}, " - f"{self._complex_dml_type(type_.value_type, implicit_bind=implicit_bind)})" + key_type = self._complex_dml_type(type_.key_type, require_precision=require_precision) + value_type = self._complex_dml_type( + type_.value_type, require_precision=require_precision ) + return f"MAP({key_type}, {value_type})" if isinstance(type_, AthenaStruct): fields = ", ".join( f"{self.preparer.quote(name)} " - f"{self._complex_dml_type(field_type, implicit_bind=implicit_bind)}" + f"{self._complex_dml_type(field_type, require_precision=require_precision)}" for name, field_type in type_.fields.items() ) return f"ROW({fields})" @@ -674,9 +690,9 @@ def _complex_dml_type(self, type_, *, implicit_bind=False): return "DOUBLE" if isinstance(type_, types.Float): return "REAL" - if implicit_bind and isinstance(type_, types.Numeric) and type_.precision is None: + if require_precision and isinstance(type_, types.Numeric) and type_.precision is None: raise exc.CompileError( - "ARRAY decimal binds require explicit Numeric precision; " + "ARRAY decimal values require explicit Numeric precision; " "specify precision and scale to avoid implicit rounding" ) return self.dialect.type_compiler_instance.process(type_) diff --git a/tests/pyathena/sqlalchemy/test_array.py b/tests/pyathena/sqlalchemy/test_array.py index e52498c1..22a4624d 100644 --- a/tests/pyathena/sqlalchemy/test_array.py +++ b/tests/pyathena/sqlalchemy/test_array.py @@ -18,17 +18,21 @@ bindparam, cast, column, + func, literal, literal_column, select, text, types, + update, ) from sqlalchemy import exc as sa_exc +from sqlalchemy.orm import declarative_base from sqlalchemy.sql import sqltypes import pyathena from pyathena.formatter import DefaultParameterFormatter +from pyathena.sqlalchemy.array import _ArrayWriteIndexType from pyathena.sqlalchemy.base import AthenaDialect from pyathena.sqlalchemy.types import ( ARRAY, @@ -83,6 +87,26 @@ def process_result_value(self, value, dialect): return tuple(value) if value is not None else None +class PrefixString(types.TypeDecorator): + impl = types.String + cache_ok = True + + def process_bind_param(self, value, dialect): + return f"prefix:{value}" + + def bind_expression(self, bindvalue): + return func.upper(bindvalue) + + +def _array_update_table(type_=None): + return Table( + "arrays", + MetaData(), + Column("id", Integer), + Column("items", type_ or AthenaArray(Integer)), + ) + + class TestAthenaArray: def test_creation_with_default(self): array_type = AthenaArray() @@ -659,3 +683,216 @@ def test_array_rewrite_requires_labels_for_literal_expressions(self, expression) .compile(dialect=AthenaDialect()) ) assert sql.startswith("SELECT anon_1.value, json_format(") + + +class TestArrayAssignmentType: + def test_binary_element_assignment_uses_native_hex_parameter(self): + table = _array_update_table(AthenaArray(types.BINARY)) + compiled = ( + table.update() + .values({table.c["items"][1]: b"\x00\xff"}) + .compile(dialect=AthenaDialect()) + ) + params = { + name: compiled._bind_processors.get(name, lambda value: value)(value) + for name, value in compiled.params.items() + } + assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params) + + def test_explicit_assignment_type_and_callable_value(self): + table = _array_update_table(AthenaArray(types.String)) + stmt = table.update().values( + {table.c["items"][1]: bindparam("value", type_=PrefixString(), callable_=lambda: "a")} + ) + compiled = stmt.compile(dialect=AthenaDialect()) + assert compiled._bind_processors["value"]("a") == "prefix:a" + assert "upper(%(value)s)" in str(compiled) + assert compiled.params["value"] == "a" + + +class TestArrayWriteIndexType: + @pytest.mark.parametrize("processor_name", ["bind_processor", "literal_processor"]) + @pytest.mark.parametrize("value", [None, True, 1.5, "1"]) + def test_rejects_non_integer_values(self, processor_name, value): + processor = getattr(_ArrayWriteIndexType(), processor_name)(AthenaDialect()) + with pytest.raises(ValueError, match="non-NULL integers"): + processor(value) + + def test_bind_and_literal_processors(self): + type_ = _ArrayWriteIndexType() + assert type_.bind_processor(AthenaDialect())(2) == 2 + assert type_.literal_processor(AthenaDialect())(2) == "2" + + +class TestArrayUpdate: + def test_multiple_updates_to_one_array_are_rejected(self): + table = _array_update_table() + values = table.c["items"] + for assignments in ( + {values[1]: 2, values[2]: 3}, + {values: [], values[1]: 2}, + {values[1]: 2, "items": []}, + ): + with pytest.raises(sa_exc.CompileError, match="one assignment"): + table.update().values(assignments).compile(dialect=AthenaDialect()) + + def test_only_final_index_may_be_a_slice(self): + table = _array_update_table(AthenaArray(Integer, dimensions=2)) + with pytest.raises(sa_exc.CompileError, match="final"): + table.update().values({table.c["items"][1:2][1]: [2]}).compile(dialect=AthenaDialect()) + + def test_bound_indices_and_values_are_reused_without_mutation(self): + table = _array_update_table() + expression = table.c["items"][bindparam("index")] + statement = table.update().values({expression: bindparam("value"), table.c.id: 2}) + compiled = statement.compile(dialect=AthenaDialect()) + assert set(compiled.params) == {"index", "value", "id"} + assert str(statement.compile(dialect=AthenaDialect())) == str(compiled) + + def test_partial_update_requires_target_table_column(self): + table = _array_update_table() + for column_ in (Column("items", AthenaArray(Integer)), _array_update_table().c["items"]): + with pytest.raises(sa_exc.CompileError, match="target table"): + table.update().values({column_[1]: 2}).compile(dialect=AthenaDialect()) + + def test_ordered_partial_update_with_sql_expression(self): + table = _array_update_table() + items = table.c["items"] + statement = table.update().ordered_values((items[2], items[1] + 1), (table.c.id, 2)) + compiled = str(statement.compile(dialect=AthenaDialect())) + assert compiled.index("SET items=") < compiled.index(", id=") + assert "element_at(arrays.items" in compiled + + def test_orm_partial_update_and_renamed_attribute_conflicts(self): + base = declarative_base() + + class Model(base): + __tablename__ = "arrays" + id = Column(Integer, primary_key=True) + values = Column("stored", AthenaArray(Integer), key="db_key") + + sql = str(update(Model).values({Model.values[1]: 2}).compile(dialect=AthenaDialect())) + assert "UPDATE arrays SET stored=concat(" in sql + for whole in (Model.values, "values"): + with pytest.raises(sa_exc.CompileError, match="one assignment"): + update(Model).values({Model.values[1]: 2, whole: []}).compile( + dialect=AthenaDialect() + ) + + +class TestArrayUpdateCompiler: + def test_bound_index_uses_write_index_processor(self): + table = _array_update_table() + compiled = ( + table.update() + .values({table.c["items"][bindparam("index")]: 9}) + .compile(dialect=AthenaDialect()) + ) + assert compiled._bind_processors["index"](2) == 2 + for value in (1.5, True, None): + with pytest.raises(ValueError, match="non-NULL integers"): + compiled._bind_processors["index"](value) + + def test_callable_index_and_slice_value(self): + table = _array_update_table(AthenaArray(types.String)) + compiled = ( + table.update() + .values({table.c["items"][bindparam("index", callable_=lambda: 1)]: "a"}) + .compile(dialect=AthenaDialect()) + ) + assert compiled.params["index"] == 1 + table.update().values( + {table.c["items"][1:2]: bindparam("values", callable_=lambda: ["a"])} + ).compile(dialect=AthenaDialect()) + + @pytest.mark.parametrize( + ("target", "value"), + [(1, 2), (4, None), (slice(2, 3), [4]), (slice(2, 2), []), (slice(None), [])], + ) + def test_partial_update_compiles_to_one_whole_column_assignment(self, target, value): + table = _array_update_table() + statement = table.update().values({table.c["items"][target]: value}).where(table.c.id == 1) + original_key = statement._generate_cache_key().key + compiled = statement.compile(dialect=AthenaDialect()) + sql = str(compiled) + assert sql.startswith("UPDATE arrays SET items=") + assert "SET element_at" not in sql + assert "SELECT" not in sql + assert "WHERE arrays.id =" in sql + assert statement._generate_cache_key().key == original_key + parameters = { + name: compiled._bind_processors.get(name, lambda v: v)(value) + for name, value in compiled.params.items() + } + formatted = DefaultParameterFormatter().format(sql, parameters) + assert "ARRAY[" in formatted + + @pytest.mark.parametrize("index", [0, -1, None, 1.5, True]) + def test_invalid_partial_update_index(self, index): + table = _array_update_table() + with pytest.raises(sa_exc.CompileError, match="indices"): + table.update().values({table.c["items"][index]: 1}).compile(dialect=AthenaDialect()) + + def test_nested_and_zero_indexed_update(self): + table = _array_update_table(AthenaArray(Integer, dimensions=2, zero_indexes=True)) + statement = table.update().values({table.c["items"][0][2]: 7}) + sql = str( + statement.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True}) + ) + assert "ARRAY[concat(" in sql + assert "sequence(" not in sql + assert "IF(1 > 0, 1," in sql + assert "IF(3 > 0, 3," in sql + + @pytest.mark.parametrize( + "array_type", + [ + types.ARRAY(Integer), + TupleArray(), + AthenaArray(Integer).with_variant(AthenaArray(Integer), "awsathena"), + ], + ) + @pytest.mark.parametrize("index", [1, bindparam("index")]) + def test_array_implementations_partial_update(self, array_type, index): + table = _array_update_table(array_type) + sql = str( + table.update().values({table.c["items"][index]: 2}).compile(dialect=AthenaDialect()) + ) + assert "SET items=concat(" in sql + + @pytest.mark.parametrize("target", [1, slice(1, 2)]) + @pytest.mark.parametrize("expression", [False, True]) + def test_decimal_assignment_requires_precision(self, target, expression): + value = [Decimal("1.23")] if isinstance(target, slice) else Decimal("1.23") + table = _array_update_table(AthenaArray(types.Numeric())) + if expression: + value = table.c["items"][target] + with pytest.raises(sa_exc.CompileError, match="precision"): + table.update().values({table.c["items"][target]: value}).compile( + dialect=AthenaDialect() + ) + table = _array_update_table(AthenaArray(types.Numeric(8, 2))) + if expression: + value = table.c["items"][target] + sql = str( + table.update() + .values({table.c["items"][target]: value}) + .compile(dialect=AthenaDialect()) + ) + assert "DECIMAL(8, 2)" in sql + + def test_write_index_expression_keeps_its_argument_types(self): + table = _array_update_table() + index = func.length("abc") + statement = table.update().values({table.c["items"][index]: 9}) + compiled = statement.compile(dialect=AthenaDialect()) + params = { + name: compiled._bind_processors.get(name, lambda value: value)(value) + for name, value in compiled.params.items() + } + assert "length('abc')" in DefaultParameterFormatter().format(str(compiled), params) + + def test_null_slice_assignment_rejected(self): + table = _array_update_table() + with pytest.raises(sa_exc.CompileError, match="non-NULL array"): + table.update().values({table.c["items"][1:2]: None}).compile(dialect=AthenaDialect()) diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index b19fe7a9..0d903d5a 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -19,10 +19,12 @@ testing, text, types, + update, ) from sqlalchemy import Table as SATable from sqlalchemy import exc as sa_exc from sqlalchemy import testing as sa_testing +from sqlalchemy.orm import Session, registry from sqlalchemy.testing import eq_, fixtures from sqlalchemy.testing.schema import Column, Table from sqlalchemy.testing.suite import * # noqa: F403 @@ -126,6 +128,217 @@ def process_result_value(self, value, dialect): return tuple(value) if value is not None else None +class ArrayUpdateTest(fixtures.TestBase): + __backend__ = True + __requires__ = ("array_type",) + + def test_element_resize_and_null(self, connection, metadata): + table = Table( + "array_element_updates", + metadata, + Column("id", Integer), + Column("items", AthenaArray(Integer)), + Column("marker", Integer), + ) + table.create(connection) + connection.execute( + table.insert(), + [ + {"id": 1, "items": [1, 2, 3]}, + {"id": 2, "items": []}, + {"id": 3, "items": None}, + ], + ) + items = table.c["items"] + connection.execute( + table.update().where(table.c.id == 1).values({items[5]: 9, table.c.marker: 42}) + ) + connection.execute(table.update().where(table.c.id == 1).values({items[2]: None})) + connection.execute(table.update().where(table.c.id > 1).values({items[2]: 7})) + eq_( + connection.execute(select(items, table.c.marker).order_by(table.c.id)).all(), + [([1, None, 3, None, 9], 42), ([None, 7], None), ([None, 7], None)], + ) + + def test_slice_resize(self, connection, metadata): + cases = [ + ([1, 2, 3], slice(2, 2), [8, 9], [1, 8, 9, 3]), + ([1, 2, 3], slice(2, 3), [8], [1, 8]), + ([1, 2, 3], slice(2, 3), [], [1]), + ([1, 2, 3], slice(3, 1), [8], [1, 2, 8, 3]), + ([1, 2, 3], slice(5, 9), [8], [1, 2, 3, None, 8]), + ([1, 2, 3], slice(2, 9), [8], [1, 8]), + ([1, 2, 3], slice(None), [8, 9], [8, 9]), + ([], slice(1, 2), [], []), + (None, slice(2, 2), [8], [None, 8]), + ] + table = Table( + "array_slice_updates", + metadata, + Column("id", Integer), + Column("items", AthenaArray(Integer)), + ) + table.create(connection) + connection.execute( + table.insert(), + [{"id": i, "items": before} for i, (before, _, _, _) in enumerate(cases)], + ) + for i, (_, bounds, replacement, _) in enumerate(cases): + connection.execute( + table.update() + .where(table.c.id == i) + .values({table.c["items"][bounds]: replacement}) + ) + eq_( + connection.execute(select(table.c["items"]).order_by(table.c.id)).scalars().all(), + [expected for _, _, _, expected in cases], + ) + + def test_nested_zero_indexed_and_cached_bindings(self, connection, metadata): + table = Table( + "array_nested_updates", + metadata, + Column("id", Integer), + Column("items", AthenaArray(Integer, dimensions=2, zero_indexes=True)), + ) + table.create(connection) + connection.execute( + table.insert(), [{"id": 1, "items": [[1], None]}, {"id": 2, "items": None}] + ) + items = table.c["items"] + statement = ( + table.update() + .where(table.c.id == bindparam("row_id")) + .values({items[bindparam("outer")][bindparam("inner")]: bindparam("value")}) + ) + connection.execute(statement, {"row_id": 1, "outer": 1, "inner": 1, "value": 7}) + connection.execute(statement, {"row_id": 2, "outer": 0, "inner": 0, "value": 8}) + connection.execute(table.update().where(table.c.id == 1).values({items[0][:0]: [4, 5]})) + eq_( + connection.execute(select(items).order_by(table.c.id)).scalars().all(), + [[[4, 5], [None, 7]], [[8]]], + ) + with pytest.raises(sa_exc.DBAPIError): + connection.execute(statement, {"row_id": 1, "outer": -1, "inner": 0, "value": 9}) + + def test_expression_values_and_indices(self, connection, metadata): + table = Table( + "array_expression_updates", + metadata, + Column("id", Integer), + Column("items", AthenaArray(Integer)), + Column("binary_items", AthenaArray(types.BINARY)), + Column("tuple_items", _ArrayTuple()), + Column("decimal_items", AthenaArray(types.Numeric(8, 2))), + Column("timestamp_items", AthenaArray(types.TIMESTAMP)), + ) + table.create(connection) + connection.execute( + table.insert().values( + id=1, + items=[1, 2, 3], + binary_items=[b"abc"], + tuple_items=[1, 2], + decimal_items=[Decimal("0.00")], + timestamp_items=[ + _datetime(2024, 1, 1, microsecond=123000), + _datetime(2024, 1, 2, microsecond=456000), + ], + ) + ) + items = table.c["items"] + connection.execute(table.update().values({table.c.decimal_items[1]: Decimal("1.23")})) + eq_(connection.execute(select(table.c.decimal_items)).scalar_one(), [Decimal("1.23")]) + connection.execute( + table.update().ordered_values( + (items[func.length("abc")], items[1] + 8), + (table.c.binary_items[1], b"\x00\xff"), + (table.c.tuple_items[bindparam("tuple_index")], bindparam("tuple_value")), + (table.c.decimal_items[1], table.c.decimal_items[1] + Decimal("3.33")), + (table.c.timestamp_items[1], table.c.timestamp_items[2]), + (table.c.id, 2), + ), + {"tuple_index": 1, "tuple_value": 5}, + ) + eq_(connection.execute(select(table.c.tuple_items)).scalar_one(), (5, 2)) + connection.execute( + table.update().values( + {items[1:2]: items[2:3].concat([4]), table.c.tuple_items[1:1]: [6, 7]} + ) + ) + eq_( + connection.execute(select(table)).one(), + ( + 2, + [2, 9, 4, 9], + [b"\x00\xff"], + (6, 7, 2), + [Decimal("4.56")], + [_datetime(2024, 1, 2, microsecond=456000)] * 2, + ), + ) + + def test_orm_and_long_array_update(self, connection, metadata): + table = Table( + "array_orm_updates", + metadata, + Column("id", Integer, primary_key=True), + Column("items", AthenaArray(Integer)), + ) + table.create(connection) + connection.execute( + table.insert().values( + id=1, + items=func.concat(func.sequence(1, 10000), literal([10001], AthenaArray(Integer))), + ) + ) + mapping = registry() + + class Record: + pass + + mapping.map_imperatively(Record, table) + try: + with Session(bind=connection) as session: + session.execute(update(Record).where(Record.id == 1).values({Record.items[1]: 99})) + session.flush() + row = connection.execute(select(table.c["items"])).scalar_one() + eq_((len(row), row[0], row[-1]), (10001, 99, 10001)) + connection.execute( + table.update().values({table.c["items"][2]: select(literal(77)).scalar_subquery()}) + ) + eq_(connection.execute(select(table.c["items"])).scalar_one()[:3], [99, 77, 3]) + finally: + mapping.dispose() + + def test_cached_literal_assignment_failures(self, connection, metadata): + table = Table( + "array_cached_invalid_updates", metadata, Column("items", AthenaArray(Integer)) + ) + table.create(connection) + connection.execute(table.insert().values(items=[1, 2])) + connection = connection.execution_options(compiled_cache={}) + items = table.c["items"] + connection.execute(table.update().values({items[1]: 3})) + with pytest.raises(sa_exc.DBAPIError, match="Invalid ARRAY index"): + connection.execute(table.update().values({items[0]: 4})) + eq_(connection.execute(select(items)).scalar_one(), [3, 2]) + connection.execute(table.update().values({items[1:2]: [7]})) + with pytest.raises(sa_exc.DBAPIError, match="NULL ARRAY slice assignment"): + connection.execute(table.update().values({items[1:2]: None})) + eq_(connection.execute(select(items)).scalar_one(), [7]) + + def test_null_slice_binding_rejected(self, connection, metadata): + table = Table("array_null_slice", metadata, Column("items", types.ARRAY(Integer))) + table.create(connection) + connection.execute(table.insert().values(items=[1, 2])) + statement = table.update().values({table.c["items"][1:2]: bindparam("replacement")}) + connection.execute(statement, {"replacement": [3]}) + with pytest.raises(sa_exc.DBAPIError): + connection.execute(statement, {"replacement": None}) + eq_(connection.execute(select(table.c["items"])).scalar_one(), [3]) + + class ArrayExpressionTest(fixtures.TestBase): __backend__ = True __requires__ = ("array_type",) From b43504faf8f24fcd5c828499f06ac427fcfffe62 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:36:45 +0900 Subject: [PATCH 08/20] Backport #811: Rerun tests once on Athena service-side query failures (cherry picked from commit a2412cc0f7383c0406fe9e848e49d9948ee622d7) Conflict resolution: benchmarks/uv.lock does not exist on 3.x: it belongs to the benchmark uv workspace from #813, which is not backported. #811's one-line change to it is dropped. docs/testing.md on 3.x predates master's rewrite in #788, so #811's two sentences are added after the 'Run test' commands, with 'The just test recipes' as the subject instead of 'They also'. justfile, pyproject.toml, and uv.lock apply unchanged, and 'uv lock --check' passes. Co-Authored-By: Claude Opus 5.5 --- docs/testing.md | 3 +++ justfile | 8 +++++--- pyproject.toml | 1 + uv.lock | 15 +++++++++++++++ 4 files changed, 24 insertions(+), 3 deletions(-) diff --git a/docs/testing.md b/docs/testing.md index e46b1eda..d21e1aae 100644 --- a/docs/testing.md +++ b/docs/testing.md @@ -40,6 +40,9 @@ $ just test sqla $ just test sqla-async ``` +The `just test` recipes rerun a failed test once when its failure message is an Athena internal error or `Invalid S3 request`, and report the first attempt's traceback. +Failures in the session setup hooks are not rerun, and a direct `pytest` invocation does not rerun. + ## Run test multiple Python versions ```bash diff --git a/justfile b/justfile index e223f613..43dddaaa 100644 --- a/justfile +++ b/justfile @@ -1,5 +1,7 @@ RUFF_VERSION := "0.14.14" TOX_VERSION := "4.34.1" +# Rerun a test once when Athena fails a query with a service-side error (#804). +PYTEST_RERUN := "--reruns 1 --rerun-show-tracebacks --only-rerun 'Amazon Athena experienced an internal error' --only-rerun 'Invalid S3 request'" # List available recipes default: @@ -34,13 +36,13 @@ _test-help: @echo " sqla-async Run SQLAlchemy async dialect tests" _test-pyathena: lint - uv run pytest -n 8 --cov pyathena --cov-report html --cov-report term tests/pyathena/ + uv run pytest -n 8 {{PYTEST_RERUN}} --cov pyathena --cov-report html --cov-report term tests/pyathena/ _test-sqla: - uv run pytest -n 8 --cov pyathena --cov-report html --cov-report term tests/sqlalchemy/ + uv run pytest -n 8 {{PYTEST_RERUN}} --cov pyathena --cov-report html --cov-report term tests/sqlalchemy/ _test-sqla-async: - uv run pytest -n 8 --cov pyathena --cov-report html --cov-report term tests/sqlalchemy/ --dburi async + uv run pytest -n 8 {{PYTEST_RERUN}} --cov pyathena --cov-report html --cov-report term tests/sqlalchemy/ --dburi async # Run tests across multiple Python versions with tox tox: diff --git a/pyproject.toml b/pyproject.toml index da33dd1e..47c77eae 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -82,6 +82,7 @@ dev = [ "pytest-asyncio", "pytest-xdist", "pytest-dependency", + "pytest-rerunfailures", "sphinx", "sphinx-multiversion", "sphinxext-opengraph", diff --git a/uv.lock b/uv.lock index 71b570e3..88c6672b 100644 --- a/uv.lock +++ b/uv.lock @@ -985,6 +985,7 @@ dev = [ { name = "pytest-asyncio" }, { name = "pytest-cov" }, { name = "pytest-dependency" }, + { name = "pytest-rerunfailures" }, { name = "pytest-xdist" }, { name = "sphinx", version = "7.4.7", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, { name = "sphinx", version = "8.2.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, @@ -1030,6 +1031,7 @@ dev = [ { name = "pytest-asyncio" }, { name = "pytest-cov" }, { name = "pytest-dependency" }, + { name = "pytest-rerunfailures" }, { name = "pytest-xdist" }, { name = "sphinx" }, { name = "sphinx-design" }, @@ -1103,6 +1105,19 @@ dependencies = [ ] sdist = { url = "https://files.pythonhosted.org/packages/7e/3b/317cc04e77d707d338540ca67b619df8f247f3f4c9f40e67bf5ea503ad94/pytest-dependency-0.6.0.tar.gz", hash = "sha256:934b0e6a39d95995062c193f7eaeed8a8ffa06ff1bcef4b62b0dc74a708bacc1", size = 19499, upload-time = "2023-12-31T20:38:54.991Z" } +[[package]] +name = "pytest-rerunfailures" +version = "16.7" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "packaging" }, + { name = "pytest" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d1/b0/6b5337f9d59b26b0069ea3d5e863c31dc04b69e51bbfb1cf2ea6328fba87/pytest_rerunfailures-16.7.tar.gz", hash = "sha256:6956ddfb65ca1d07e7e3d99e2c5359d82300f4cb062e6049563bd2c106f72d5c", size = 59666, upload-time = "2026-09-17T07:08:48.871Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/c1/d3/07ea35102cf2020ddaaac368d2a6f0bcc63d6523fb8918805bac3eb8b98f/pytest_rerunfailures-16.7-py3-none-any.whl", hash = "sha256:edf1886209c2b7dafe35b5bf1708d6ec40ccf6c6b357f0f02807efcec0204c99", size = 23952, upload-time = "2026-09-17T07:08:47.635Z" }, +] + [[package]] name = "pytest-xdist" version = "3.6.1" From 2290413394152446fb002c6dc9971b1b94ea446e Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:37:26 +0900 Subject: [PATCH 09/20] Backport #815: Give each test session its own S3 Tables namespace (cherry picked from commit 00828a310e6f07aa300f7a2156c7c5c7af5f768f) Conflict resolution: Dropped #815's changes to files that do not exist on 3.x: .github/workflows/database-sweep.yaml, scripts/sweep_databases.py, and scripts/tests/test_sweep_databases.py come from #779 (the master-only weekly sweep), and tests/pyathena/test_glue.py comes from #803 (excluded). docs/testing.md on 3.x predates master's rewrite and has no Amazon S3 Tables section, so #815's update to that section is dropped. In tests/pyathena/conftest.py, the import block adds only #815's 'import functools'; master's 'import uuid' there came from #768, which is not backported. Co-Authored-By: Claude Opus 5.5 --- .github/workflows/test-suite.yaml | 8 +- .../cloudformation/github_actions_oidc.yaml | 10 -- tests/__init__.py | 14 +-- tests/pyathena/conftest.py | 95 +++++++++++++++++-- tests/pyathena/sqlalchemy/test_base.py | 28 +++--- 5 files changed, 115 insertions(+), 40 deletions(-) diff --git a/.github/workflows/test-suite.yaml b/.github/workflows/test-suite.yaml index 4edf52b6..9a4669de 100644 --- a/.github/workflows/test-suite.yaml +++ b/.github/workflows/test-suite.yaml @@ -20,11 +20,11 @@ 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. + # Registered S3 Tables catalog (s3tablescatalog/) for the S3 + # Tables tests; each test session creates and deletes its own namespace. + # The table bucket and the AWS analytics-services integration are + # provisioned out of band. AWS_ATHENA_S3_TABLES_CATALOG: s3tablescatalog/laughingman7743-pyathena-s3-tables - AWS_ATHENA_S3_TABLES_NAMESPACE: pyathena strategy: fail-fast: false diff --git a/scripts/cloudformation/github_actions_oidc.yaml b/scripts/cloudformation/github_actions_oidc.yaml index 0212439f..e17ae0e8 100644 --- a/scripts/cloudformation/github_actions_oidc.yaml +++ b/scripts/cloudformation/github_actions_oidc.yaml @@ -24,10 +24,6 @@ Parameters: Type: String Default: "laughingman7743-pyathena-s3-tables" Description: Name of the Amazon S3 Tables table bucket used by the SQLAlchemy S3 Tables tests. - S3TablesNamespaceName: - Type: String - Default: "pyathena" - Description: Namespace created in the S3 Tables table bucket for the SQLAlchemy S3 Tables tests. OIDCProviderArn: Type: String Default: "" @@ -353,12 +349,6 @@ Resources: UnreferencedDays: 1 NoncurrentDays: 1 - S3TablesNamespace: - Type: AWS::S3Tables::Namespace - Properties: - TableBucketARN: !GetAtt S3TablesBucket.TableBucketARN - Namespace: !Sub "${S3TablesNamespaceName}" - # One-time, per-Region integration: register S3 Tables into the Glue Data # Catalog as the federated `s3tablescatalog` catalog so Athena can query them. # IAM_ALLOWED_PRINCIPALS default permissions select IAM-based access control diff --git a/tests/__init__.py b/tests/__init__.py index be0927b8..124fe300 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -31,16 +31,16 @@ def __init__(self): ) self.default_work_group = os.getenv("AWS_ATHENA_DEFAULT_WORKGROUP", "primary") self.managed_work_group = os.getenv("AWS_ATHENA_MANAGED_WORKGROUP") - # Optional Amazon S3 Tables configuration for the SQLAlchemy dialect tests. - # `s3tables_catalog` is the registered table-bucket catalog, e.g. - # "s3tablescatalog/"; `s3tables_namespace` is a namespace that - # already exists in it. These are optional so the suite stays runnable - # without an S3 Tables bucket; the S3 Tables tests skip when they are unset. - self.s3tables_catalog = os.getenv("AWS_ATHENA_S3_TABLES_CATALOG") - self.s3tables_namespace = os.getenv("AWS_ATHENA_S3_TABLES_NAMESPACE") self.schema = "pyathena_test_" + "".join( random.choices(string.ascii_lowercase + string.digits, k=10) ) + # Optional Amazon S3 Tables configuration. `s3tables_catalog` is the + # registered table-bucket catalog, e.g. "s3tablescatalog/". + # The test session creates its own namespace in it, named like the schema, + # because a table dropped during a listing of a shared namespace fails + # the listing. The S3 Tables tests skip when the catalog is unset. + self.s3tables_catalog = os.getenv("AWS_ATHENA_S3_TABLES_CATALOG") + self.s3tables_namespace = self.schema if self.s3tables_catalog else None self.s3_filesystem_test_file_key = ( f"{self.s3_staging_key}{self.schema}/filesystem/test_read/test.dat" ) diff --git a/tests/pyathena/conftest.py b/tests/pyathena/conftest.py index 7e0f249f..ca236f0c 100644 --- a/tests/pyathena/conftest.py +++ b/tests/pyathena/conftest.py @@ -1,4 +1,5 @@ import contextlib +import functools from io import BytesIO from pathlib import Path @@ -12,16 +13,96 @@ def pytest_sessionstart(session): - _upload_rows() - with contextlib.closing(connect()) as conn, conn.cursor() as cursor: - _create_database(cursor) - _create_table(cursor) + # pytest skips pytest_sessionfinish after a failed pytest_sessionstart, so + # a failure after the namespace is created deletes it here. + is_test_process = _is_test_process(session.config) + if is_test_process: + _create_s3tables_namespace() + try: + _upload_rows() + with contextlib.closing(connect()) as conn, conn.cursor() as cursor: + _create_database(cursor) + _create_table(cursor) + except BaseException: + if is_test_process: + _delete_s3tables_namespace() + raise def pytest_sessionfinish(session): - with contextlib.closing(connect()) as conn, conn.cursor() as cursor: - _drop_database(cursor) - _delete_rows() + # Each cleanup step runs even if an earlier one fails. + try: + with contextlib.closing(connect()) as conn, conn.cursor() as cursor: + _drop_database(cursor) + finally: + try: + _delete_rows() + finally: + if _is_test_process(session.config): + _delete_s3tables_namespace() + + +def _is_test_process(config): + """Whether this process runs tests, rather than only controlling xdist workers. + + Args: + config: The pytest config. + + Returns: + False for the pytest-xdist controller, True for a worker or a run without + workers. + """ + return hasattr(config, "workerinput") or not getattr(config.option, "numprocesses", None) + + +@functools.cache +def _s3tables(): + """Return an S3 Tables client and the ARN of ``ENV.s3tables_catalog``'s table bucket. + + The ARN uses the client's region, so the two always agree. + + Returns: + The client and the table bucket's ARN. + + Raises: + ValueError: If ``AWS_ATHENA_S3_TABLES_CATALOG`` is not + ``s3tablescatalog/``. + """ + prefix, _, bucket = ENV.s3tables_catalog.partition("/") + if prefix != "s3tablescatalog" or not bucket: + raise ValueError( + "AWS_ATHENA_S3_TABLES_CATALOG must be s3tablescatalog/, " + f"not {ENV.s3tables_catalog!r}." + ) + client = boto3.client("s3tables") + account = boto3.client("sts").get_caller_identity()["Account"] + region = client.meta.region_name + return client, f"arn:aws:s3tables:{region}:{account}:bucket/{bucket}" + + +def _create_s3tables_namespace(): + """Create this process's S3 Tables namespace when S3 Tables are configured.""" + if not ENV.s3tables_catalog: + return + client, arn = _s3tables() + client.create_namespace(tableBucketARN=arn, namespace=[ENV.s3tables_namespace]) + + +def _delete_s3tables_namespace(): + """Delete this process's S3 Tables namespace and any table left in it.""" + if not ENV.s3tables_catalog: + return + client, arn = _s3tables() + tables = [ + table["name"] + for page in client.get_paginator("list_tables").paginate( + tableBucketARN=arn, namespace=ENV.s3tables_namespace + ) + for table in page["tables"] + ] + for table in tables: + client.delete_table(tableBucketARN=arn, namespace=ENV.s3tables_namespace, name=table) + client.delete_namespace(tableBucketARN=arn, namespace=ENV.s3tables_namespace) def _upload_rows(): diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 989ee191..e85d183b 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -26,22 +26,26 @@ ) from tests.pyathena.conftest import ENV -# Amazon S3 Tables tests need a pre-provisioned table-bucket catalog and namespace. -# Skip them unless AWS_ATHENA_S3_TABLES_CATALOG / AWS_ATHENA_S3_TABLES_NAMESPACE are set. +# Amazon S3 Tables tests need a pre-provisioned table-bucket catalog; the session +# creates its own namespace in it. +# Skip them unless AWS_ATHENA_S3_TABLES_CATALOG is set. requires_s3_tables = pytest.mark.skipif( - not ENV.s3tables_catalog or not ENV.s3tables_namespace, - reason="AWS_ATHENA_S3_TABLES_CATALOG / AWS_ATHENA_S3_TABLES_NAMESPACE are not configured", + not ENV.s3tables_catalog, + reason="AWS_ATHENA_S3_TABLES_CATALOG is not configured", ) def unique_s3tables_table_name(base: str) -> str: - """Return a per-run-unique S3 Tables table name. + """Return a unique S3 Tables table name. - Other integration tests isolate themselves with a random per-process - ``ENV.schema``, but the S3 Tables tests share one fixed namespace - (``ENV.s3tables_namespace``). The CI matrix runs ``tests/pyathena`` once per - Python version in parallel against the same account, so a fixed table name - would collide across those concurrent jobs; a random suffix keeps them apart. + The session's namespace (``ENV.s3tables_namespace``) is its own, but a + rerun of a failed test would find the table an earlier attempt left there. + + Args: + base: The name to extend. + + Returns: + ``base`` with a random suffix. """ return f"{base}_{uuid.uuid4().hex[:8]}" @@ -2166,8 +2170,8 @@ def test_create_s3tables_iceberg_table(self, engine): assert tblproperties["table_type"] == "ICEBERG" finally: # Idempotent unquoted drop: tolerates a table that was never created - # while still surfacing systematic DROP failures, which would - # otherwise leak tables into the shared fixed namespace. + # while still surfacing systematic DROP failures; the session's + # namespace cleanup removes anything left behind. conn.execute(text(f"DROP TABLE IF EXISTS {schema}.{table_name}")) @requires_s3_tables From 29a4d0aa08bd65fc075b52bc4f11e19a06483883 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:37:37 +0900 Subject: [PATCH 10/20] Backport #814: Run the three test suites in parallel (cherry picked from commit 76ccd107104770a194ed4ed34ca59a42d62e4518) Co-Authored-By: Claude Opus 5.5 --- .github/workflows/test.yaml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/test.yaml b/.github/workflows/test.yaml index d8e6a3ee..1a23c4fe 100644 --- a/.github/workflows/test.yaml +++ b/.github/workflows/test.yaml @@ -22,6 +22,8 @@ concurrency: group: ${{ github.workflow }}-${{ github.ref }} cancel-in-progress: true +# The three suites create their own schemas and tables, so they run in +# parallel; each is still a separate job for "Re-run failed jobs". jobs: test: uses: ./.github/workflows/test-suite.yaml @@ -29,13 +31,11 @@ jobs: test-type: pyathena test-sqla: - needs: [test] uses: ./.github/workflows/test-suite.yaml with: test-type: sqla test-sqla-async: - needs: [test-sqla] uses: ./.github/workflows/test-suite.yaml with: test-type: sqla_async From c5211802b356ed55ab849caa8e1de1a657e8a339 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:38:45 +0900 Subject: [PATCH 11/20] Backport #825: Re-enable SQLAlchemy parameter and SQL rendering compliance tests (cherry picked from commit 69009827ab6e849095e2569dcb2e6b0e7471be29) Conflict resolution: Both test files conflicted with neighboring changes from #771, #777, and #818 (not backported). Applied only #825's own changes, and the added and removed lines match #825's diff exactly. In test_compiler.py: the union and DefaultParameterFormatter imports, DIFFICULT_PARAMETER_NAMES, _format_sql, and the four rendering tests at the end of TestAthenaStatementCompiler. In test_suite.py: import DifficultParametersTest and drop its del, add the DifficultParametersTest subclass before InsertBehaviorTest, and remove the test_limit_render_multiple_times and StringTest skips. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/sqlalchemy/test_compiler.py | 71 ++++++++++++++++++++++ tests/sqlalchemy/test_suite.py | 29 ++++----- 2 files changed, 86 insertions(+), 14 deletions(-) diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index a160cfc1..a1b8c99d 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -21,18 +21,45 @@ table, text, types, + union, ) from sqlalchemy.engine.url import make_url from sqlalchemy.sql import literal, literal_column, operators from sqlalchemy.sql.compiler import FROM_LINTING from sqlalchemy.sql.ddl import CreateTable +from pyathena.formatter import DefaultParameterFormatter from pyathena.sqlalchemy.base import AthenaDialect from pyathena.sqlalchemy.compiler import AthenaTypeCompiler from pyathena.sqlalchemy.pandas import AthenaPandasDialect from pyathena.sqlalchemy.types import ARRAY, MAP, STRUCT, AthenaArray, AthenaMap, AthenaStruct from tests import ENV +# Bind parameter names from SQLAlchemy's DifficultParametersTest. +DIFFICULT_PARAMETER_NAMES = [ + "boring", + "per cent", + "per % cent", + "%percent", + "par(ens)", + "percent%(ens)yah", + "col:ons", + "_starts_with_underscore", + "dot.s", + "more :: %colons%", + "_name", + "___name", + "[BracketsAndCase]", + "42numbers", + "percent%signs", + "has spaces", + "/slashes/", + "more/slashes", + "q?marks", + "1param", + "1col:on", +] + class TestAthenaTypeCompiler: def test_visit_struct_empty(self): @@ -414,6 +441,50 @@ def test_stepped_array_aggregate_with_inferred_element_type(self): assert "Unsupported ARRAY slice step" in sql assert "ARRAY(NULL)" not in sql + def _format_sql(self, statement, parameters=None): + """Format a statement with the parameter names SQLAlchemy sends to the cursor. + + Bind processors are not applied, so use this only for types without one. + + Args: + statement: SQLAlchemy statement to compile. + parameters: Values for the statement's bind parameters. + + Returns: + The SQL string produced by ``DefaultParameterFormatter``. + """ + compiled = statement.compile(dialect=self.dialect) + # Expand before escaping names, as SQLAlchemy's execution context does. + expanded = compiled.construct_expanded_state(parameters, escape_names=False) + escaped_names = compiled.escaped_bind_names + formatted_params = {escaped_names.get(k, k): v for k, v in expanded.parameters.items()} + return DefaultParameterFormatter().format(expanded.statement, formatted_params) + + @pytest.mark.parametrize("name", DIFFICULT_PARAMETER_NAMES) + def test_difficult_bind_parameter_name(self, name): + id_ = column("id", Integer) + stmt = select(id_).where(id_ == bindparam(name, type_=Integer)) + assert self._format_sql(stmt, {name: 3}).endswith("WHERE id = 3") + + @pytest.mark.parametrize("name", DIFFICULT_PARAMETER_NAMES) + def test_difficult_expanding_bind_parameter_name(self, name): + id_ = column("id", Integer) + stmt = select(id_).where(id_.in_(bindparam(name, value=[1, 2]))) + assert self._format_sql(stmt, {name: [4, 1]}).endswith("WHERE id IN (4, 1)") + + def test_limit_rendered_multiple_times(self): + limited = ( + select(self.test_table.c.id).order_by(self.test_table.c.id).limit(1).scalar_subquery() + ) + sql = self._format_sql(union(select(limited), select(limited)).subquery().select()) + assert sql.count("LIMIT 1") == 2 + assert "%(" not in sql + + @pytest.mark.parametrize("pattern", ["%B%", "A%C", "A%C%Z", "%(x)s"]) + def test_like_pattern_is_not_truncated(self, pattern): + stmt = select(column("x", String)).where(column("x", String).like(pattern)) + assert self._format_sql(stmt).endswith(f"WHERE x LIKE '{pattern}'") + class TestAthenaDDLCompiler: """Compile-only (no AWS) tests for the DDL compiler's S3 Tables support. diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 0d903d5a..2f6aa737 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -30,12 +30,12 @@ from sqlalchemy.testing.suite import * # noqa: F403 from sqlalchemy.testing.suite import BinaryTest as _BinaryTest from sqlalchemy.testing.suite import CTETest as _CTETest +from sqlalchemy.testing.suite import DifficultParametersTest as _DifficultParametersTest from sqlalchemy.testing.suite import FetchLimitOffsetTest as _FetchLimitOffsetTest from sqlalchemy.testing.suite import HasTableTest as _HasTableTest from sqlalchemy.testing.suite import InsertBehaviorTest as _InsertBehaviorTest from sqlalchemy.testing.suite import IntegerTest as _IntegerTest from sqlalchemy.testing.suite import SimpleUpdateDeleteTest as _SimpleUpdateDeleteTest -from sqlalchemy.testing.suite import StringTest as _StringTest from pyathena.sqlalchemy.types import ( AthenaArray, @@ -49,7 +49,6 @@ del ComponentReflectionTestExtra # noqa: F821 del CompositeKeyReflectionTest # noqa: F821 del DateTimeMicrosecondsTest # noqa: F821 -del DifficultParametersTest # noqa: F821 del DistinctOnTest # noqa: F821 del HasIndexTest # noqa: F821 del IdentityAutoincrementTest # noqa: F821 @@ -729,6 +728,20 @@ def test_has_table_cache(self, metadata): pass +class DifficultParametersTest(_DifficultParametersTest): + # Each case creates an Iceberg table, and the formatter inlines bind values before the + # query reaches Athena. test_round_trip_same_named_column still sends every name live + # as a column name and an explicit bind parameter; test_compiler.py covers explicit + # expanding binds offline, because round-trip IN parameters get sanitized names. + @pytest.mark.skip("Explicit binds are covered live by test_round_trip_same_named_column.") + def test_standalone_bindparam_escape(self, paramname, connection, multirow_fixture): + pass + + @pytest.mark.skip("Expanding binds are covered offline by the dialect compile tests.") + def test_standalone_bindparam_escape_expanding(self, paramname, connection, multirow_fixture): + pass + + class InsertBehaviorTest(_InsertBehaviorTest): @pytest.mark.skip("Athena does not support auto-incrementing.") def test_insert_from_select_autoinc(self, connection): @@ -765,21 +778,9 @@ def test_expr_limit_simple_offset(self, connection): def test_expr_offset(self, connection): pass - @pytest.mark.skip("TODO") - def test_limit_render_multiple_times(self, connection): - # TODO - pass - class IntegerTest(_IntegerTest): @pytest.mark.skip("TODO") def test_huge_int(self, integer_round_trip, intvalue): # TODO pass - - -class StringTest(_StringTest): - @pytest.mark.skip("TODO") - def test_dont_truncate_rightside(self, metadata, connection, expr, expected): - # TODO - pass From c4dfbbd1003f09e9446fd1b9855598f71b3a3dc5 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:39:03 +0900 Subject: [PATCH 12/20] Backport #826: SQLAlchemy compliance: audit numeric range and float precision (cherry picked from commit 3ddb868f85c0a168253ca15fc5286f8970c7392a) Conflict resolution: tests/sqlalchemy/test_suite.py conflicted with import lines from #771 (not backported). Applied only #826's own change there, removing the IntegerTest import and its test_huge_int skip, so the removed lines match #826's diff exactly. Co-Authored-By: Claude Opus 5.5 --- docs/sqlalchemy.md | 11 +++++++++++ pyathena/sqlalchemy/compiler.py | 2 +- pyathena/sqlalchemy/requirements.py | 12 +++++++++--- tests/pyathena/sqlalchemy/test_compiler.py | 16 ++++++++++++++++ tests/sqlalchemy/test_suite.py | 8 -------- 5 files changed, 37 insertions(+), 12 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index e3229472..0fd9796c 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -663,6 +663,17 @@ engine_arrow = create_engine( ) ``` +## Floating-point types + +| SQLAlchemy type | Table DDL | CAST | +|---|---|---| +| `Float`, `FLOAT`, `REAL` | `FLOAT` | `REAL` | +| `Double`, `DOUBLE`, `DOUBLE_PRECISION` | `DOUBLE` | `DOUBLE` | + +Athena `FLOAT` and `REAL` are the same 32-bit floating-point type, which keeps about seven significant digits. +Use `Double` (SQLAlchemy 2.0+) for 64-bit values. +`Float(precision)` does not change the Athena type. + ## Complex data types ### STRUCT type support diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 3ca46bfa..6c0bad22 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -645,7 +645,7 @@ def visit_cast(self, cast: Cast[Any], **kwargs): type_clause = "CHAR" elif isinstance(cast.type, (types.LargeBinary, 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/pyathena/sqlalchemy/requirements.py b/pyathena/sqlalchemy/requirements.py index 12f22d53..dfa9c71a 100644 --- a/pyathena/sqlalchemy/requirements.py +++ b/pyathena/sqlalchemy/requirements.py @@ -89,14 +89,20 @@ def timestamp_microseconds(self): @property def precision_generic_float_type(self): - # TODO: AssertionError: - # {Decimal('15.7563820'), Decimal('15.7563830')} != {Decimal('15.7563827')} - return unsupported() + return exclusions.skip_if( + lambda _: True, + "Generic Float maps to Athena REAL, a 32-bit type with about seven " + "significant digits; use Double for 64-bit values.", + ) @property def precision_numerics_many_significant_digits(self): return supported() + @property + def precision_numerics_retains_significant_digits(self): + return supported() + @property def window_functions(self): return supported() diff --git a/tests/pyathena/sqlalchemy/test_compiler.py b/tests/pyathena/sqlalchemy/test_compiler.py index a1b8c99d..df0e6b5b 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -173,6 +173,22 @@ def test_visit_json(self): result = compiler.visit_JSON(json_type) assert result == "JSON" + @pytest.mark.parametrize( + ("type_", "ddl", "cast_type"), + [ + (types.Float(), "FLOAT", "REAL"), + (types.FLOAT(), "FLOAT", "REAL"), + (types.REAL(), "FLOAT", "REAL"), + (types.Double(), "DOUBLE", "DOUBLE"), + (types.DOUBLE(), "DOUBLE", "DOUBLE"), + (types.DOUBLE_PRECISION(), "DOUBLE", "DOUBLE"), + ], + ) + def test_floating_point_types(self, type_, ddl, cast_type): + dialect = AthenaDialect() + 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.""" diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 2f6aa737..c44ed24e 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -34,7 +34,6 @@ from sqlalchemy.testing.suite import FetchLimitOffsetTest as _FetchLimitOffsetTest from sqlalchemy.testing.suite import HasTableTest as _HasTableTest from sqlalchemy.testing.suite import InsertBehaviorTest as _InsertBehaviorTest -from sqlalchemy.testing.suite import IntegerTest as _IntegerTest from sqlalchemy.testing.suite import SimpleUpdateDeleteTest as _SimpleUpdateDeleteTest from pyathena.sqlalchemy.types import ( @@ -777,10 +776,3 @@ def test_expr_limit_simple_offset(self, connection): @pytest.mark.skip("Athena does not support expressions in the offset clause.") def test_expr_offset(self, connection): pass - - -class IntegerTest(_IntegerTest): - @pytest.mark.skip("TODO") - def test_huge_int(self, integer_round_trip, intvalue): - # TODO - pass From 70a763d0bcfc83e70fea1e78948f0c2538641bdc Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:40:05 +0900 Subject: [PATCH 13/20] Backport #837: Run AWS test suites only on ready pull requests and related changes (cherry picked from commit e104a42fda4b2e890a343b85d67c47b0aac8e6c4) Conflict resolution: Dropped #837's changes to files and passages that do not exist on 3.x: .agents/skills/development-workflow/SKILL.md, and the development-workflow line in AGENTS.md (both from #778). In .github/workflows/test-suite.yaml, applied only #837's skip-spark and skip-sqla inputs and PYTEST_ADDOPTS; the neighboring max-parallel input and SDK retry settings came from #777 (not backported). docs/testing.md on 3.x predates master's rewrite, so #837's CI policy text is added unchanged at the start of its GitHub Actions section. test.yaml and pyproject.toml apply unchanged. Co-Authored-By: Claude Opus 5.5 --- .github/workflows/test-suite.yaml | 11 ++++ .github/workflows/test.yaml | 88 +++++++++++++++++++++++++++++-- docs/testing.md | 22 ++++++++ pyproject.toml | 1 + 4 files changed, 117 insertions(+), 5 deletions(-) diff --git a/.github/workflows/test-suite.yaml b/.github/workflows/test-suite.yaml index 9a4669de..3e42509d 100644 --- a/.github/workflows/test-suite.yaml +++ b/.github/workflows/test-suite.yaml @@ -6,6 +6,14 @@ on: test-type: required: true type: string + skip-spark: + description: Skip the Spark tests of the PyAthena suite + type: boolean + default: false + skip-sqla: + description: Skip the SQLAlchemy tests of the PyAthena suite + type: boolean + default: false jobs: run: @@ -15,6 +23,9 @@ jobs: env: TEST_TYPE: ${{ inputs.test-type }} + PYTEST_ADDOPTS: >- + ${{ inputs.skip-spark && '--ignore=tests/pyathena/spark --ignore=tests/pyathena/aio/spark' || '' }} + ${{ inputs.skip-sqla && '--ignore=tests/pyathena/sqlalchemy --ignore=tests/pyathena/aio/sqlalchemy' || '' }} AWS_DEFAULT_REGION: us-west-2 AWS_ATHENA_S3_STAGING_DIR: s3://laughingman7743-pyathena/github/ AWS_ATHENA_WORKGROUP: pyathena diff --git a/.github/workflows/test.yaml b/.github/workflows/test.yaml index 1a23c4fe..ca728a42 100644 --- a/.github/workflows/test.yaml +++ b/.github/workflows/test.yaml @@ -2,14 +2,21 @@ 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. paths-ignore applies to + # both, so a pull request that now changes only docs starts neither. + types: [opened, synchronize, reopened, ready_for_review, converted_to_draft] paths-ignore: - 'docs/**' - '**.md' + # The scheduled run executes every suite, including the ones that pull + # requests only run when related files change. 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: before a release, on demand for + # a pull request, and to refresh the README status badge after a transient + # failure on the default branch. workflow_dispatch: permissions: @@ -22,20 +29,91 @@ concurrency: group: ${{ github.workflow }}-${{ github.ref }} cancel-in-progress: true -# The three suites create their own schemas and tables, so they run in -# parallel; each is still a separate job for "Re-run failed jobs". 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 AWS suites. Draft and external-fork pull requests run none. + # A ready pull request always runs the PyAthena suite; it runs the + # SQLAlchemy tests (the compliance suites and the PyAthena suite's + # SQLAlchemy tests) and the Spark tests only when their code, tests, + # dependencies, or this workflow change. + changes: + 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: + pull-requests: read + outputs: + sqla: ${{ steps.filter.outputs.sqla }} + spark: ${{ steps.filter.outputs.spark }} + steps: + - id: filter + env: + GH_TOKEN: ${{ github.token }} + EVENT_NAME: ${{ github.event_name }} + REPO: ${{ github.repository }} + PR_NUMBER: ${{ github.event.pull_request.number }} + run: | + if [[ "$EVENT_NAME" != "pull_request" ]]; then + echo "sqla=true" >> "$GITHUB_OUTPUT" + echo "spark=true" >> "$GITHUB_OUTPUT" + exit 0 + fi + files=$(gh api "repos/$REPO/pulls/$PR_NUMBER/files" --paginate --jq '.[].filename') + printf 'Changed files:\n%s\n' "$files" + shared='^(\.github/workflows/test(-suite)?\.yaml|justfile|pyproject\.toml|uv\.lock)$' + sqla="$shared|^pyathena/(aio/)?sqlalchemy/|^tests/sqlalchemy/|^tests/pyathena/(aio/)?sqlalchemy/|^setup\.cfg$" + spark="$shared|^pyathena/(aio/)?spark/|^tests/pyathena/(aio/)?spark/" + if grep -qE "$sqla" <<< "$files"; then + echo "sqla=true" >> "$GITHUB_OUTPUT" + else + echo "sqla=false" >> "$GITHUB_OUTPUT" + fi + if grep -qE "$spark" <<< "$files"; then + echo "spark=true" >> "$GITHUB_OUTPUT" + else + echo "spark=false" >> "$GITHUB_OUTPUT" + fi + + # The three suites create their own schemas and tables, so they run in + # parallel; each is still a separate job for "Re-run failed jobs". test: + needs: changes uses: ./.github/workflows/test-suite.yaml with: test-type: pyathena + skip-spark: ${{ needs.changes.outputs.spark != 'true' }} + skip-sqla: ${{ needs.changes.outputs.sqla != 'true' }} test-sqla: + needs: changes + if: needs.changes.outputs.sqla == 'true' uses: ./.github/workflows/test-suite.yaml with: test-type: sqla test-sqla-async: + needs: changes + if: needs.changes.outputs.sqla == 'true' uses: ./.github/workflows/test-suite.yaml with: test-type: sqla_async diff --git a/docs/testing.md b/docs/testing.md index d21e1aae..88654d1b 100644 --- a/docs/testing.md +++ b/docs/testing.md @@ -68,6 +68,28 @@ $ just lint ## GitHub Actions +The Test workflow runs for pull requests that change files other than `docs/` and Markdown. +It runs the offline checks (`just lint`) on each of them, including Drafts and external forks, and runs the AWS suites as follows: + +| Trigger | PyAthena suite | SQLAlchemy tests | Spark tests | +| --- | --- | --- | --- | +| Draft pull request | No | No | No | +| Ready pull request from a branch of this repository | Yes | When related files change | When related files change | +| Weekly schedule and manual dispatch | Yes | Yes | Yes | + +The SQLAlchemy tests are the compliance suites and the PyAthena suite's `tests/pyathena/sqlalchemy/` and `tests/pyathena/aio/sqlalchemy/`. +The Spark tests are the PyAthena suite's `tests/pyathena/spark/` and `tests/pyathena/aio/spark/`. +When the SQLAlchemy or Spark tests do not run, the PyAthena suite runs without them. +For the SQLAlchemy tests, the related files are `pyathena/sqlalchemy/`, `pyathena/aio/sqlalchemy/`, `tests/sqlalchemy/`, their PyAthena suite test directories, and `setup.cfg`. +For the Spark tests, they are `pyathena/spark/`, `pyathena/aio/spark/`, and their PyAthena suite test directories. +Changes to `pyproject.toml`, `uv.lock`, `justfile`, or the Test workflows run both. +For a pull request from a branch of this repository that still changes files other than `docs/` and Markdown, marking the Draft ready for review starts the AWS jobs, and converting it back to Draft cancels AWS jobs still running. +To run every suite on a branch, dispatch the workflow: + +```bash +gh workflow run test.yaml --ref +``` + GitHub Actions uses OpenID Connect (OIDC) to access AWS resources. You will need to refer to the [GitHub Actions documentation](https://docs.github.com/actions/deployment/security-hardening-your-deployments/configuring-openid-connect-in-amazon-web-services) to configure it. The CloudFormation templates for creating GitHub OIDC Provider and IAM Role can be found in the [aws-actions/configure-aws-credentials repository](https://github.com/aws-actions/configure-aws-credentials#sample-iam-role-cloudformation-template). diff --git a/pyproject.toml b/pyproject.toml index 47c77eae..de188806 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -229,6 +229,7 @@ commands = sqla_async: just test sqla-async passenv = TOXENV + PYTEST_ADDOPTS AWS_* GITHUB_* """ From 91c922b96eff817189d2fae14d1be9b3e96a5a2a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:41:15 +0900 Subject: [PATCH 14/20] Backport #863: Run pull-request AWS tests on the newest Python version only (cherry picked from commit 3cb9663de659d7cb8cb6e049a017c43e2b428b3a) Conflict resolution: .github/workflows/test-suite.yaml conflicted with the max-parallel input from #777 (not backported). Applied only #863's own changes: the python-versions input and the matrix taken from it. test.yaml and docs/testing.md apply unchanged. Co-Authored-By: Claude Opus 5.5 --- .github/workflows/test-suite.yaml | 6 +++++- .github/workflows/test.yaml | 24 ++++++++++++++++++------ docs/testing.md | 12 ++++++------ 3 files changed, 29 insertions(+), 13 deletions(-) diff --git a/.github/workflows/test-suite.yaml b/.github/workflows/test-suite.yaml index 3e42509d..84231a3c 100644 --- a/.github/workflows/test-suite.yaml +++ b/.github/workflows/test-suite.yaml @@ -6,6 +6,10 @@ on: test-type: required: true type: string + python-versions: + description: JSON array of the Python versions to test + required: true + type: string skip-spark: description: Skip the Spark tests of the PyAthena suite type: boolean @@ -40,7 +44,7 @@ jobs: 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 ca728a42..075aa659 100644 --- a/.github/workflows/test.yaml +++ b/.github/workflows/test.yaml @@ -49,11 +49,12 @@ jobs: tool: just - run: just lint - # Selects the AWS suites. Draft and external-fork pull requests run none. - # A ready pull request always runs the PyAthena suite; it runs the - # SQLAlchemy tests (the compliance suites and the PyAthena suite's + # Selects the AWS suites and Python versions. Draft and external-fork pull + # requests run none. A ready pull request always runs the PyAthena suite; it + # runs the SQLAlchemy tests (the compliance suites and the PyAthena suite's # SQLAlchemy tests) and the Spark tests only when their code, tests, - # dependencies, or this workflow change. + # dependencies, or this workflow change. Pull requests test the newest + # Python version only; the schedule and dispatch test every version. changes: if: >- github.event_name != 'pull_request' || @@ -65,6 +66,7 @@ jobs: outputs: sqla: ${{ steps.filter.outputs.sqla }} spark: ${{ steps.filter.outputs.spark }} + python-versions: ${{ steps.filter.outputs.python-versions }} steps: - id: filter env: @@ -72,12 +74,19 @@ jobs: EVENT_NAME: ${{ github.event_name }} REPO: ${{ github.repository }} PR_NUMBER: ${{ github.event.pull_request.number }} + # 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: | if [[ "$EVENT_NAME" != "pull_request" ]]; then - echo "sqla=true" >> "$GITHUB_OUTPUT" - echo "spark=true" >> "$GITHUB_OUTPUT" + { + echo "python-versions=$(jq -c '.' <<< "$PYTHON_VERSIONS")" + echo "sqla=true" + echo "spark=true" + } >> "$GITHUB_OUTPUT" exit 0 fi + echo "python-versions=$(jq -c '[last]' <<< "$PYTHON_VERSIONS")" >> "$GITHUB_OUTPUT" files=$(gh api "repos/$REPO/pulls/$PR_NUMBER/files" --paginate --jq '.[].filename') printf 'Changed files:\n%s\n' "$files" shared='^(\.github/workflows/test(-suite)?\.yaml|justfile|pyproject\.toml|uv\.lock)$' @@ -101,6 +110,7 @@ jobs: uses: ./.github/workflows/test-suite.yaml with: test-type: pyathena + python-versions: ${{ needs.changes.outputs.python-versions }} skip-spark: ${{ needs.changes.outputs.spark != 'true' }} skip-sqla: ${{ needs.changes.outputs.sqla != 'true' }} @@ -110,6 +120,7 @@ jobs: uses: ./.github/workflows/test-suite.yaml with: test-type: sqla + python-versions: ${{ needs.changes.outputs.python-versions }} test-sqla-async: needs: changes @@ -117,3 +128,4 @@ jobs: uses: ./.github/workflows/test-suite.yaml with: test-type: sqla_async + python-versions: ${{ needs.changes.outputs.python-versions }} diff --git a/docs/testing.md b/docs/testing.md index 88654d1b..b29376e4 100644 --- a/docs/testing.md +++ b/docs/testing.md @@ -71,11 +71,11 @@ $ just lint The Test workflow runs for pull requests that change files other than `docs/` and Markdown. It runs the offline checks (`just lint`) on each of them, including Drafts and external forks, and runs the AWS suites as follows: -| Trigger | PyAthena suite | SQLAlchemy tests | Spark tests | -| --- | --- | --- | --- | -| Draft pull request | No | No | No | -| Ready pull request from a branch of this repository | Yes | When related files change | When related files change | -| Weekly schedule and manual dispatch | Yes | Yes | Yes | +| Trigger | PyAthena suite | SQLAlchemy tests | Spark tests | Python versions | +| --- | --- | --- | --- | --- | +| Draft pull request | No | No | No | None | +| Ready pull request from a branch of this repository | Yes | When related files change | When related files change | Newest supported | +| Weekly schedule and manual dispatch | Yes | Yes | Yes | All supported | The SQLAlchemy tests are the compliance suites and the PyAthena suite's `tests/pyathena/sqlalchemy/` and `tests/pyathena/aio/sqlalchemy/`. The Spark tests are the PyAthena suite's `tests/pyathena/spark/` and `tests/pyathena/aio/spark/`. @@ -84,7 +84,7 @@ For the SQLAlchemy tests, the related files are `pyathena/sqlalchemy/`, `pyathen For the Spark tests, they are `pyathena/spark/`, `pyathena/aio/spark/`, and their PyAthena suite test directories. Changes to `pyproject.toml`, `uv.lock`, `justfile`, or the Test workflows run both. For a pull request from a branch of this repository that still changes files other than `docs/` and Markdown, marking the Draft ready for review starts the AWS jobs, and converting it back to Draft cancels AWS jobs still running. -To run every suite on a branch, dispatch the workflow: +To run every suite on every supported Python version on a branch, dispatch the workflow: ```bash gh workflow run test.yaml --ref From 8340b4e07477aba3f62024313c730e4f3a00c077 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:42:11 +0900 Subject: [PATCH 15/20] Backport #862: Run the full test matrix in the Release workflow before publishing (cherry picked from commit 9a1a72d7f25b4372b5f1c66ec190fea132f60228) Conflict resolution: Both conflicts came from lines outside #862. .github/workflows/docs-trigger.yaml on master carries the license header from #805 (not backported), so only #862's rename, comment, workflow_run trigger, and success condition are applied. In docs/testing.md, #862's two sentences on the Release workflow are added before 3.x's OIDC paragraph; master's neighboring fork-policy text came from its #788 rewrite. release.yaml, test.yaml, and docs.yaml apply unchanged. Co-Authored-By: Claude Opus 5.5 --- .github/workflows/docs-trigger.yaml | 10 ++++--- .github/workflows/docs.yaml | 17 ++++++++++++ .github/workflows/release.yaml | 17 +++++++++--- .github/workflows/test.yaml | 42 +++++++++++++++++++++++------ docs/testing.md | 10 +++++-- 5 files changed, 80 insertions(+), 16 deletions(-) diff --git a/.github/workflows/docs-trigger.yaml b/.github/workflows/docs-trigger.yaml index 6bd3ca82..5016a665 100644 --- a/.github/workflows/docs-trigger.yaml +++ b/.github/workflows/docs-trigger.yaml @@ -1,14 +1,18 @@ -name: Trigger Docs on Tag +name: Trigger Docs on Release +# Rebuilds the documentation after the Release workflow succeeds, so a tag it +# refuses to publish does not trigger a documentation build. on: - push: - tags: ['v*'] + workflow_run: + workflows: [Release] + types: [completed] permissions: actions: write jobs: trigger-docs: + if: github.event.workflow_run.conclusion == 'success' runs-on: ubuntu-latest steps: - uses: actions/github-script@ed597411d8f924073f98dfc5c65a23a2325f34cd # v8.0.0 diff --git a/.github/workflows/docs.yaml b/.github/workflows/docs.yaml index 2439aaf8..a0fa3f78 100644 --- a/.github/workflows/docs.yaml +++ b/.github/workflows/docs.yaml @@ -24,6 +24,23 @@ jobs: uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 with: fetch-depth: 0 # Fetch all history for sphinx-multiversion + # sphinx-multiversion builds every version tag in the checkout. A tag + # gets its GitHub release only after the Release workflow's tests and + # PyPI upload succeed, so drop tags without one: a release still in + # progress or refused by its tests is not documented. + - name: Drop unreleased version tags + env: + GH_TOKEN: ${{ github.token }} + run: | + released=$(gh api --paginate "repos/$GITHUB_REPOSITORY/releases?per_page=100" \ + --jq '.[] | select(.draft | not) | .tag_name') + if [[ -z "$released" ]]; then + echo "::error::No published GitHub releases found" + exit 1 + fi + git tag --list 'v*' | while read -r tag; do + grep -qxF "$tag" <<< "$released" || git tag --delete "$tag" + done - name: Setup Pages uses: actions/configure-pages@45bfe0192ca1faeb007ade9deae92b16b8254a0d # v6.0.0 - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7.6.0 diff --git a/.github/workflows/release.yaml b/.github/workflows/release.yaml index d7918824..54f2e3e1 100644 --- a/.github/workflows/release.yaml +++ b/.github/workflows/release.yaml @@ -5,13 +5,24 @@ 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 + pull-requests: read + release: + needs: test runs-on: ubuntu-latest + permissions: + id-token: write + contents: write env: PYTHON_VERSION: '3.12' diff --git a/.github/workflows/test.yaml b/.github/workflows/test.yaml index 075aa659..4f2cf48a 100644 --- a/.github/workflows/test.yaml +++ b/.github/workflows/test.yaml @@ -11,13 +11,21 @@ on: - 'docs/**' - '**.md' # The scheduled run executes every suite, including the ones that pull - # requests only run when related files change. + # requests only run when related files change, on the newest Python version. schedule: - cron: '0 0 * * 0' - # Runs every suite on the selected branch: before a release, on demand for - # a pull request, and to refresh the README status badge after a transient - # failure on the default branch. + # Runs every suite on the selected branch: on demand for a pull request, and + # to refresh the README status badge after a transient failure on the + # default branch. 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 @@ -53,8 +61,9 @@ jobs: # requests run none. A ready pull request always runs the PyAthena suite; it # runs the SQLAlchemy tests (the compliance suites and the PyAthena suite's # SQLAlchemy tests) and the Spark tests only when their code, tests, - # dependencies, or this workflow change. Pull requests test the newest - # Python version only; the schedule and dispatch test every version. + # dependencies, or this workflow change. Pull requests and the schedule test + # the newest Python version; a dispatch tests the requested versions or + # every version, and the Release workflow every version. changes: if: >- github.event_name != 'pull_request' || @@ -74,19 +83,36 @@ jobs: EVENT_NAME: ${{ github.event_name }} REPO: ${{ github.repository }} PR_NUMBER: ${{ github.event.pull_request.number }} + 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 | schedule) + 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" if [[ "$EVENT_NAME" != "pull_request" ]]; then { - echo "python-versions=$(jq -c '.' <<< "$PYTHON_VERSIONS")" echo "sqla=true" echo "spark=true" } >> "$GITHUB_OUTPUT" exit 0 fi - echo "python-versions=$(jq -c '[last]' <<< "$PYTHON_VERSIONS")" >> "$GITHUB_OUTPUT" files=$(gh api "repos/$REPO/pulls/$PR_NUMBER/files" --paginate --jq '.[].filename') printf 'Changed files:\n%s\n' "$files" shared='^(\.github/workflows/test(-suite)?\.yaml|justfile|pyproject\.toml|uv\.lock)$' diff --git a/docs/testing.md b/docs/testing.md index b29376e4..56c53f32 100644 --- a/docs/testing.md +++ b/docs/testing.md @@ -75,7 +75,9 @@ It runs the offline checks (`just lint`) on each of them, including Drafts and e | --- | --- | --- | --- | --- | | Draft pull request | No | No | No | None | | Ready pull request from a branch of this repository | Yes | When related files change | When related files change | Newest supported | -| Weekly schedule and manual dispatch | Yes | Yes | Yes | All supported | +| Weekly schedule | Yes | Yes | Yes | Newest supported | +| Manual dispatch | Yes | Yes | Yes | Requested, or all supported | +| Release tag (Release workflow) | Yes | Yes | Yes | All supported | The SQLAlchemy tests are the compliance suites and the PyAthena suite's `tests/pyathena/sqlalchemy/` and `tests/pyathena/aio/sqlalchemy/`. The Spark tests are the PyAthena suite's `tests/pyathena/spark/` and `tests/pyathena/aio/spark/`. @@ -84,12 +86,16 @@ For the SQLAlchemy tests, the related files are `pyathena/sqlalchemy/`, `pyathen For the Spark tests, they are `pyathena/spark/`, `pyathena/aio/spark/`, and their PyAthena suite test directories. Changes to `pyproject.toml`, `uv.lock`, `justfile`, or the Test workflows run both. For a pull request from a branch of this repository that still changes files other than `docs/` and Markdown, marking the Draft ready for review starts the AWS jobs, and converting it back to Draft cancels AWS jobs still running. -To run every suite on every supported Python version on a branch, dispatch the workflow: +To run every suite on a branch, dispatch the workflow; it tests every supported Python version unless `python-versions` lists some of them: ```bash gh workflow run test.yaml --ref +gh workflow run test.yaml --ref -f python-versions=3.11,3.14 ``` +The Release workflow runs the same suites on every supported Python version for the tagged commit before building, and publishes nothing unless all of them pass. +If they fail, nothing is published, and the documentation leaves out the tag because it only lists tags with a GitHub release; delete the tag, fix the failure, and push the tag again. + GitHub Actions uses OpenID Connect (OIDC) to access AWS resources. You will need to refer to the [GitHub Actions documentation](https://docs.github.com/actions/deployment/security-hardening-your-deployments/configuring-openid-connect-in-amazon-web-services) to configure it. The CloudFormation templates for creating GitHub OIDC Provider and IAM Role can be found in the [aws-actions/configure-aws-credentials repository](https://github.com/aws-actions/configure-aws-credentials#sample-iam-role-cloudformation-template). From caaae8fba1d69b96c12b428d84af24203c4f6785 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:42:44 +0900 Subject: [PATCH 16/20] Backport #866: Set up the PyAthena test fixtures only in processes that run tests (cherry picked from commit 569fa2dddcb0ec5260aa98b0e9c058bf9f652603) Conflict resolution: tests/resources/queries/create_database.sql.jinja2 and insert_into_table.sql.jinja2 conflicted only with the license headers from #805 (not backported). Applied #866's own changes: removed the DROP DATABASE line from create_database.sql.jinja2 and deleted the unused insert_into_table.sql.jinja2, which nothing on 3.x references either. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/conftest.py | 15 ++++++++------- .../queries/create_database.sql.jinja2 | 1 - .../resources/queries/create_table.sql.jinja2 | 16 ---------------- .../queries/insert_into_table.sql.jinja2 | 19 ------------------- 4 files changed, 8 insertions(+), 43 deletions(-) delete mode 100644 tests/resources/queries/insert_into_table.sql.jinja2 diff --git a/tests/pyathena/conftest.py b/tests/pyathena/conftest.py index ca236f0c..1ced60fe 100644 --- a/tests/pyathena/conftest.py +++ b/tests/pyathena/conftest.py @@ -13,23 +13,25 @@ def pytest_sessionstart(session): + # The pytest-xdist controller runs no tests, so it sets up nothing. + if not _is_test_process(session.config): + return + _create_s3tables_namespace() # pytest skips pytest_sessionfinish after a failed pytest_sessionstart, so # a failure after the namespace is created deletes it here. - is_test_process = _is_test_process(session.config) - if is_test_process: - _create_s3tables_namespace() try: _upload_rows() with contextlib.closing(connect()) as conn, conn.cursor() as cursor: _create_database(cursor) _create_table(cursor) except BaseException: - if is_test_process: - _delete_s3tables_namespace() + _delete_s3tables_namespace() raise def pytest_sessionfinish(session): + if not _is_test_process(session.config): + return # Each cleanup step runs even if an earlier one fails. try: with contextlib.closing(connect()) as conn, conn.cursor() as cursor: @@ -38,8 +40,7 @@ def pytest_sessionfinish(session): try: _delete_rows() finally: - if _is_test_process(session.config): - _delete_s3tables_namespace() + _delete_s3tables_namespace() def _is_test_process(config): diff --git a/tests/resources/queries/create_database.sql.jinja2 b/tests/resources/queries/create_database.sql.jinja2 index fc0efe76..2adafa15 100644 --- a/tests/resources/queries/create_database.sql.jinja2 +++ b/tests/resources/queries/create_database.sql.jinja2 @@ -1,2 +1 @@ -DROP DATABASE IF EXISTS {{ schema }} CASCADE; CREATE DATABASE IF NOT EXISTS {{ schema }}; diff --git a/tests/resources/queries/create_table.sql.jinja2 b/tests/resources/queries/create_table.sql.jinja2 index 7ca9de0f..5dc44273 100644 --- a/tests/resources/queries/create_table.sql.jinja2 +++ b/tests/resources/queries/create_table.sql.jinja2 @@ -1,4 +1,3 @@ -DROP TABLE IF EXISTS {{ schema }}.one_row; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.one_row (number_of_rows INT COMMENT 'some comment') COMMENT 'table comment' ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE @@ -6,14 +5,12 @@ LOCATION '{{ s3_staging_dir }}{{ schema }}/one_row/'; CREATE OR REPLACE VIEW {{ schema }}.view_one_row AS SELECT * FROM {{ schema }}.one_row; -DROP TABLE IF EXISTS {{ schema }}.many_rows; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.many_rows ( a INT ) ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/many_rows/'; -DROP TABLE IF EXISTS {{ schema }}.one_row_complex; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.one_row_complex ( col_boolean BOOLEAN, col_tinyint TINYINT, @@ -35,7 +32,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.one_row_complex ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/one_row_complex/'; -DROP TABLE IF EXISTS {{ schema }}.partition_table; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.partition_table ( a STRING ) @@ -43,7 +39,6 @@ PARTITIONED BY (b INT) ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/partition_table/'; -DROP TABLE IF EXISTS {{ schema }}.integer_na_values; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.integer_na_values ( a INT, b INT @@ -51,7 +46,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.integer_na_values ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/integer_na_values/'; -DROP TABLE IF EXISTS {{ schema }}.boolean_na_values; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.boolean_na_values ( a BOOLEAN, b BOOLEAN @@ -59,7 +53,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.boolean_na_values ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/boolean_na_values/'; -DROP TABLE IF EXISTS {{ schema }}.execute_many; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many ( a INT, b STRING @@ -67,7 +60,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many/'; -DROP TABLE IF EXISTS {{ schema }}.execute_many_aio; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_aio ( a INT, b STRING @@ -75,7 +67,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_aio ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_aio/'; -DROP TABLE IF EXISTS {{ schema }}.execute_many_pandas; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_pandas ( a INT, b STRING @@ -83,7 +74,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_pandas ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_pandas/'; -DROP TABLE IF EXISTS {{ schema }}.execute_many_pandas_unload_auto; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_pandas_unload_auto ( a INT, b STRING @@ -91,7 +81,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_pandas_unload_auto ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_pandas_unload_auto/'; -DROP TABLE IF EXISTS {{ schema }}.execute_many_pandas_unload_pyarrow; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_pandas_unload_pyarrow ( a INT, b STRING @@ -99,7 +88,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_pandas_unload_pyar ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_pandas_unload_pyarrow/'; -DROP TABLE IF EXISTS {{ schema }}.execute_many_arrow; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_arrow ( a INT, b STRING @@ -107,7 +95,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_arrow ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_arrow/'; -DROP TABLE IF EXISTS {{ schema }}.execute_many_arrow_unload; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_arrow_unload ( a INT, b STRING @@ -115,7 +102,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_arrow_unload ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_arrow_unload/'; -DROP TABLE IF EXISTS {{ schema }}.execute_many_polars; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_polars ( a INT, b STRING @@ -123,7 +109,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_polars ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_polars/'; -DROP TABLE IF EXISTS {{ schema }}.execute_many_polars_unload; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_polars_unload ( a INT, b STRING @@ -131,7 +116,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_polars_unload ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_polars_unload/'; -DROP TABLE IF EXISTS {{ schema }}.parquet_with_compression; CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.parquet_with_compression ( a INT ) diff --git a/tests/resources/queries/insert_into_table.sql.jinja2 b/tests/resources/queries/insert_into_table.sql.jinja2 deleted file mode 100644 index d3989551..00000000 --- a/tests/resources/queries/insert_into_table.sql.jinja2 +++ /dev/null @@ -1,19 +0,0 @@ -INSERT INTO {{ schema } -SELECT - true, - CAST(127 AS TINYINT), - CAST(32767 AS SMALLINT), - 2147483647, - 9223372036854775807, - CAST(0.5 AS REAL), - 0.25, - 'a string', - CAST('varchar' AS VARCHAR(10)), - CAST('2017-01-01 00:00:00' AS TIMESTAMP), - CAST('2017-01-02' AS DATE), - CAST('123' AS VARBINARY), - ARRAY[1, 2], - MAP(ARRAY[1, 3], ARRAY[2, 4]), - CAST(ROW(1, 2) AS ROW(a INTEGER, b INTEGER)), - CAST(DECIMAL '0.1' AS DECIMAL(10, 1)) -FROM (select 1) dual; From f0380bfb578a87d8cc2b133a4baed4af3ca2294d Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:43:14 +0900 Subject: [PATCH 17/20] Backport #867: Insert SQLAlchemy compliance fixture rows once for read-only classes (cherry picked from commit 645e04ebf9f818dfcc5336bd8af66adf1587f33a) Conflict resolution: tests/sqlalchemy/test_suite.py conflicted with import lines from #771 (not backported). Applied only #867's own changes, and the added and removed lines match #867's diff exactly: the ten suite imports, the _InsertFixtureRowsOnce mixin, the mixin on FetchLimitOffsetTest, and the ten read-only subclasses at the end. Co-Authored-By: Claude Opus 5.5 --- tests/sqlalchemy/test_suite.py | 81 +++++++++++++++++++++++++++++++++- 1 file changed, 80 insertions(+), 1 deletion(-) diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index c44ed24e..b572f574 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -29,12 +29,22 @@ from sqlalchemy.testing.schema import Column, Table from sqlalchemy.testing.suite import * # noqa: F403 from sqlalchemy.testing.suite import BinaryTest as _BinaryTest +from sqlalchemy.testing.suite import CollateTest as _CollateTest +from sqlalchemy.testing.suite import CompoundSelectTest as _CompoundSelectTest from sqlalchemy.testing.suite import CTETest as _CTETest +from sqlalchemy.testing.suite import DeprecatedCompoundSelectTest as _DeprecatedCompoundSelectTest from sqlalchemy.testing.suite import DifficultParametersTest as _DifficultParametersTest +from sqlalchemy.testing.suite import ExistsTest as _ExistsTest +from sqlalchemy.testing.suite import ExpandingBoundInTest as _ExpandingBoundInTest from sqlalchemy.testing.suite import FetchLimitOffsetTest as _FetchLimitOffsetTest from sqlalchemy.testing.suite import HasTableTest as _HasTableTest from sqlalchemy.testing.suite import InsertBehaviorTest as _InsertBehaviorTest +from sqlalchemy.testing.suite import OrderByLabelTest as _OrderByLabelTest +from sqlalchemy.testing.suite import PostCompileParamsTest as _PostCompileParamsTest +from sqlalchemy.testing.suite import RowFetchTest as _RowFetchTest +from sqlalchemy.testing.suite import SameNamedSchemaTableTest as _SameNamedSchemaTableTest from sqlalchemy.testing.suite import SimpleUpdateDeleteTest as _SimpleUpdateDeleteTest +from sqlalchemy.testing.suite import WindowFunctionTest as _WindowFunctionTest from pyathena.sqlalchemy.types import ( AthenaArray, @@ -61,6 +71,35 @@ del UuidTest # noqa: F821 +class _InsertFixtureRowsOnce: + """Insert a compliance class's fixture rows once instead of around every test. + + SQLAlchemy's ``TablesTest`` inserts the fixture rows before each test and + deletes them after it; on Athena each of those statements is an Iceberg + commit. Mix this into classes whose tests only read their fixture rows. + """ + + run_inserts = "once" + run_deletes = None + + @classmethod + def _setup_once_inserts(cls): + """Insert the fixture rows, dropping the class's tables if that fails. + + SQLAlchemy registers the class teardown that drops the tables only after + the inserts succeed. A table left behind would break a later class on the + same worker that defines a table of the same name differently. + + Raises: + BaseException: Whatever the inserts raised, after the tables are dropped. + """ + try: + super()._setup_once_inserts() # type: ignore[misc] + except BaseException: + cls._teardown_once_metadata_bind() # type: ignore[attr-defined] + raise + + class BinaryTest(_BinaryTest): @sa_testing.combinations(types.LargeBinary, types.BINARY, types.VARBINARY, argnames="datatype") @sa_testing.combinations( @@ -756,7 +795,7 @@ def test_no_results_for_non_returning_insert(self, connection, style, executeman pass -class FetchLimitOffsetTest(_FetchLimitOffsetTest): +class FetchLimitOffsetTest(_InsertFixtureRowsOnce, _FetchLimitOffsetTest): @pytest.mark.skip("Athena does not support expressions in the offset clause.") def test_simple_limit_expr_offset(self, connection): pass @@ -776,3 +815,43 @@ def test_expr_limit_simple_offset(self, connection): @pytest.mark.skip("Athena does not support expressions in the offset clause.") def test_expr_offset(self, connection): pass + + +class CollateTest(_InsertFixtureRowsOnce, _CollateTest): + pass + + +class CompoundSelectTest(_InsertFixtureRowsOnce, _CompoundSelectTest): + pass + + +class DeprecatedCompoundSelectTest(_InsertFixtureRowsOnce, _DeprecatedCompoundSelectTest): + pass + + +class ExistsTest(_InsertFixtureRowsOnce, _ExistsTest): + pass + + +class ExpandingBoundInTest(_InsertFixtureRowsOnce, _ExpandingBoundInTest): + pass + + +class OrderByLabelTest(_InsertFixtureRowsOnce, _OrderByLabelTest): + pass + + +class PostCompileParamsTest(_InsertFixtureRowsOnce, _PostCompileParamsTest): + pass + + +class RowFetchTest(_InsertFixtureRowsOnce, _RowFetchTest): + pass + + +class SameNamedSchemaTableTest(_InsertFixtureRowsOnce, _SameNamedSchemaTableTest): + pass + + +class WindowFunctionTest(_InsertFixtureRowsOnce, _WindowFunctionTest): + pass From 08645a10e41d69351d91b9076985905256cf5ecb Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:44:13 +0900 Subject: [PATCH 18/20] Backport #873: Give the executemany and partition tests their own tables (cherry picked from commit c783d8b4a77aefcd648ae5a649967bb867119f6f) Conflict resolution: The conflicts came from #768 (not backported), which added the executemany_table fixture, the executemany rowcount tests, and a rowcount assertion. Applied only #873's own changes. In tests/pyathena/conftest.py, added the empty_table fixture and the uuid import it needs, which master had through #768. In test_cursor.py and aio/test_cursor.py, test_executemany uses empty_table, and 3.x's rowcount == -1 assertions are kept. test_show_partition uses its own partition table. The other cursor tests, create_table.sql.jinja2, and tests/sqlalchemy/conftest.py apply unchanged. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/aio/test_cursor.py | 6 +- tests/pyathena/arrow/test_cursor.py | 7 +-- tests/pyathena/conftest.py | 28 +++++++++ tests/pyathena/pandas/test_cursor.py | 9 +-- tests/pyathena/polars/test_cursor.py | 7 +-- tests/pyathena/test_cursor.py | 35 +++++++---- .../resources/queries/create_table.sql.jinja2 | 63 ------------------- tests/sqlalchemy/conftest.py | 5 +- 8 files changed, 66 insertions(+), 94 deletions(-) diff --git a/tests/pyathena/aio/test_cursor.py b/tests/pyathena/aio/test_cursor.py index d18a5be9..26733c90 100644 --- a/tests/pyathena/aio/test_cursor.py +++ b/tests/pyathena/aio/test_cursor.py @@ -288,14 +288,14 @@ async def test_cancel_initial(self, aio_cursor): with pytest.raises(ProgrammingError): await aio_cursor.cancel() - async def test_executemany(self, aio_cursor): + async def test_executemany(self, aio_cursor, empty_table): rows = [(1, "foo"), (2, "bar"), (3, "jim o'rourke")] await aio_cursor.executemany( - "INSERT INTO execute_many_aio (a, b) VALUES (%(a)d, %(b)s)", + f"INSERT INTO {empty_table} (a, b) VALUES (%(a)d, %(b)s)", [{"a": a, "b": b} for a, b in rows], ) assert aio_cursor.rowcount == -1 - await aio_cursor.execute("SELECT * FROM execute_many_aio") + await aio_cursor.execute(f"SELECT * FROM {empty_table}") assert sorted(await aio_cursor.fetchall()) == list(rows) async def test_executemany_fetch(self, aio_cursor): diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index 9b32854c..11a3dabb 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -725,14 +725,13 @@ def test_ctas(self, arrow_cursor): [{"cursor_kwargs": {"unload": False}}, {"cursor_kwargs": {"unload": True}}], indirect=["arrow_cursor"], ) - def test_executemany(self, arrow_cursor): + def test_executemany(self, arrow_cursor, empty_table): rows = [(1, "foo"), (2, "bar"), (3, "jim o'rourke")] - table_name = f"execute_many_arrow{'_unload' if arrow_cursor._unload else ''}" arrow_cursor.executemany( - f"INSERT INTO {table_name} (a, b) VALUES (%(a)d, %(b)s)", + f"INSERT INTO {empty_table} (a, b) VALUES (%(a)d, %(b)s)", [{"a": a, "b": b} for a, b in rows], ) - arrow_cursor.execute(f"SELECT * FROM {table_name}") + arrow_cursor.execute(f"SELECT * FROM {empty_table}") assert sorted(arrow_cursor.fetchall()) == list(rows) @pytest.mark.parametrize( diff --git a/tests/pyathena/conftest.py b/tests/pyathena/conftest.py index 1ced60fe..31063eab 100644 --- a/tests/pyathena/conftest.py +++ b/tests/pyathena/conftest.py @@ -1,5 +1,6 @@ import contextlib import functools +import uuid from io import BytesIO from pathlib import Path @@ -222,6 +223,33 @@ def cursor(request): yield from _cursor(Cursor, request) +@pytest.fixture +def empty_table(): + """Create an empty ``(a INT, b STRING)`` text table for one test and drop it on teardown. + + The table has its own connection, so a test with any cursor type, including + the aio cursors, can write to it. + + Yields: + The table name qualified with ``ENV.schema``. + """ + table_name = f"empty_{uuid.uuid4().hex}" + table = f"{ENV.schema}.{table_name}" + with contextlib.closing(connect(schema_name=ENV.schema)) as conn, conn.cursor() as cursor: + try: + cursor.execute( + f""" + CREATE EXTERNAL TABLE {table} (a INT, b STRING) + ROW FORMAT DELIMITED FIELDS TERMINATED BY '\\t' LINES TERMINATED BY '\\n' + STORED AS TEXTFILE + LOCATION '{ENV.s3_staging_dir}{ENV.schema}/{table_name}/' + """ + ) + yield table + finally: + cursor.execute(f"DROP TABLE IF EXISTS {table}") + + @pytest.fixture def dict_cursor(request): from pyathena.cursor import DictCursor diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 34aa0968..00bd26fb 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -1131,16 +1131,13 @@ def test_boolean_na_values(self, pandas_cursor, parquet_engine): ], indirect=["pandas_cursor"], ) - def test_executemany(self, pandas_cursor, parquet_engine): + def test_executemany(self, pandas_cursor, parquet_engine, empty_table): rows = [(1, "foo"), (2, "bar"), (3, "jim o'rourke")] - table_name = "execute_many_pandas" + ( - f"_unload_{parquet_engine}" if pandas_cursor._unload else "" - ) pandas_cursor.executemany( - f"INSERT INTO {table_name} (a, b) VALUES (%(a)d, %(b)s)", + f"INSERT INTO {empty_table} (a, b) VALUES (%(a)d, %(b)s)", [{"a": a, "b": b} for a, b in rows], ) - pandas_cursor.execute(f"SELECT * FROM {table_name}", engine=parquet_engine) + pandas_cursor.execute(f"SELECT * FROM {empty_table}", engine=parquet_engine) assert sorted(pandas_cursor.fetchall()) == list(rows) @pytest.mark.parametrize( diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index 6e4b482e..ab3215c9 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -411,14 +411,13 @@ def test_empty_result_unload(self, polars_cursor): [{"cursor_kwargs": {"unload": False}}, {"cursor_kwargs": {"unload": True}}], indirect=["polars_cursor"], ) - def test_executemany(self, polars_cursor): + def test_executemany(self, polars_cursor, empty_table): rows = [(1, "foo"), (2, "bar"), (3, "jim o'rourke")] - table_name = f"execute_many_polars{'_unload' if polars_cursor._unload else ''}" polars_cursor.executemany( - f"INSERT INTO {table_name} (a, b) VALUES (%(a)d, %(b)s)", + f"INSERT INTO {empty_table} (a, b) VALUES (%(a)d, %(b)s)", [{"a": a, "b": b} for a, b in rows], ) - polars_cursor.execute(f"SELECT * FROM {table_name}") + polars_cursor.execute(f"SELECT * FROM {empty_table}") assert sorted(polars_cursor.fetchall()) == list(rows) @pytest.mark.parametrize( diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index 0b48b4a8..cea162fd 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -6,6 +6,7 @@ import string import threading import time +import uuid from concurrent import futures from concurrent.futures.thread import ThreadPoolExecutor from datetime import date, datetime, timezone @@ -796,17 +797,29 @@ def test_no_ops(self): conn.close() def test_show_partition(self, cursor): - location = f"{ENV.s3_staging_dir}{ENV.schema}/partition_table/" - for i in range(10): + table_name = f"partition_{uuid.uuid4().hex}" + table = f"{ENV.schema}.{table_name}" + location = f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/" + try: cursor.execute( + f""" + CREATE EXTERNAL TABLE {table} (a STRING) + PARTITIONED BY (b INT) + LOCATION '{location}' """ - ALTER TABLE partition_table ADD PARTITION (b=%(b)d) - LOCATION %(location)s - """, - {"b": i, "location": location}, ) - cursor.execute("SHOW PARTITIONS partition_table") - assert sorted(cursor.fetchall()) == [(f"b={i}",) for i in range(10)] + for i in range(10): + cursor.execute( + f""" + ALTER TABLE {table} ADD PARTITION (b=%(b)d) + LOCATION %(location)s + """, + {"b": i, "location": location}, + ) + cursor.execute(f"SHOW PARTITIONS {table}") + assert sorted(cursor.fetchall()) == [(f"b={i}",) for i in range(10)] + finally: + cursor.execute(f"DROP TABLE IF EXISTS {table}") @pytest.mark.parametrize("cursor", [{"work_group": ENV.work_group}], indirect=["cursor"]) def test_workgroup(self, cursor): @@ -819,15 +832,15 @@ def test_no_s3_staging_dir(self, cursor): cursor.execute("SELECT * FROM one_row") assert cursor.output_location - def test_executemany(self, cursor): + def test_executemany(self, cursor, empty_table): rows = [(1, "foo"), (2, "bar"), (3, "jim o'rourke")] cursor.executemany( - "INSERT INTO execute_many (a, b) VALUES (%(a)d, %(b)s)", + f"INSERT INTO {empty_table} (a, b) VALUES (%(a)d, %(b)s)", [{"a": a, "b": b} for a, b in rows], ) # rowcount is not supported for executemany assert cursor.rowcount == -1 - cursor.execute("SELECT * FROM execute_many") + cursor.execute(f"SELECT * FROM {empty_table}") assert sorted(cursor.fetchall()) == list(rows) def test_executemany_fetch(self, cursor): diff --git a/tests/resources/queries/create_table.sql.jinja2 b/tests/resources/queries/create_table.sql.jinja2 index 5dc44273..3d69a8c5 100644 --- a/tests/resources/queries/create_table.sql.jinja2 +++ b/tests/resources/queries/create_table.sql.jinja2 @@ -53,69 +53,6 @@ CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.boolean_na_values ( ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE LOCATION '{{ s3_staging_dir }}{{ schema }}/boolean_na_values/'; -CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many ( - a INT, - b STRING -) -ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE -LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many/'; - -CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_aio ( - a INT, - b STRING -) -ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE -LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_aio/'; - -CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_pandas ( - a INT, - b STRING -) -ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE -LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_pandas/'; - -CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_pandas_unload_auto ( - a INT, - b STRING -) -ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE -LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_pandas_unload_auto/'; - -CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_pandas_unload_pyarrow ( - a INT, - b STRING -) -ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE -LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_pandas_unload_pyarrow/'; - -CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_arrow ( - a INT, - b STRING -) -ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE -LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_arrow/'; - -CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_arrow_unload ( - a INT, - b STRING -) -ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE -LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_arrow_unload/'; - -CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_polars ( - a INT, - b STRING -) -ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE -LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_polars/'; - -CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.execute_many_polars_unload ( - a INT, - b STRING -) -ROW FORMAT DELIMITED FIELDS TERMINATED BY '\t' LINES TERMINATED BY '\n' STORED AS TEXTFILE -LOCATION '{{ s3_staging_dir }}{{ schema }}/execute_many_polars_unload/'; - CREATE EXTERNAL TABLE IF NOT EXISTS {{ schema }}.parquet_with_compression ( a INT ) diff --git a/tests/sqlalchemy/conftest.py b/tests/sqlalchemy/conftest.py index b7056916..b17e9dd7 100644 --- a/tests/sqlalchemy/conftest.py +++ b/tests/sqlalchemy/conftest.py @@ -46,9 +46,8 @@ def _awsathena_create_db(cfg, eng, ident): @drop_db.for_db("awsathena") def _awsathena_drop_db(cfg, eng, ident): with eng.begin() as conn: - conn.exec_driver_sql(f"DROP DATABASE {ident} CASCADE") - conn.exec_driver_sql(f"DROP DATABASE {ident}_test_schema CASCADE") - conn.exec_driver_sql(f"DROP DATABASE {ident}_test_schema_2 CASCADE") + for name in (ident, f"{ident}_test_schema", f"{ident}_test_schema_2"): + conn.exec_driver_sql(f"DROP DATABASE IF EXISTS {name} CASCADE") @configure_follower.for_db("awsathena") From 39163451e3c0a4b53fb270df1e3d686a10d48bc6 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 21:44:50 +0900 Subject: [PATCH 19/20] Backport #870: fix: render Hive STRUCT syntax in table column DDL (cherry picked from commit cf673c2f9118b98797ddcb5e50beba7f8a70dd64) Conflict resolution: tests/pyathena/sqlalchemy/test_compiler.py conflicted because master places #870's test_complex_cast_keeps_dml_syntax among the temporal-literal and TypeDecorator CAST tests from #818 and #836 (not backported). Only #870's test is added there, and the file's added and removed lines match #870's diff exactly. compiler.py, test_base.py, and docs/sqlalchemy.md apply unchanged. Co-Authored-By: Claude Opus 5.5 --- docs/sqlalchemy.md | 20 ++- pyathena/sqlalchemy/compiler.py | 62 ++++--- tests/pyathena/sqlalchemy/test_base.py | 63 ++++++- tests/pyathena/sqlalchemy/test_compiler.py | 188 ++++++++++++++++++++- 4 files changed, 301 insertions(+), 32 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 0fd9796c..7b75c0ea 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -708,12 +708,17 @@ 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. +That includes top-level columns, fields of a STRUCT, STRUCT values inside MAP, and STRUCT values inside ARRAY. +Integer fields, and integer MAP keys and values, use `INT` in that DDL. +`CAST` and other SQL expressions keep `ROW(...)`, `MAP(...)`, and `ARRAY(...)`, and spell integers as `INTEGER`. + #### Querying STRUCT data PyAthena automatically converts STRUCT data between different formats: @@ -853,13 +858,16 @@ This generates the following SQL structure: ```sql CREATE TABLE products ( - id INTEGER, + id INT, attributes MAP, - metrics MAP, - categories MAP + metrics MAP, + categories MAP ) ``` +`CREATE TABLE` renders integer MAP keys and values as `INT`. +`CAST` still spells those integers as `INTEGER`. + #### Querying MAP data PyAthena automatically converts MAP data between different formats: diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 6c0bad22..590cf142 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -88,6 +88,9 @@ class AthenaTypeCompiler(GenericTypeCompiler): - MAP: Key-value pair collections - ARRAY: Ordered collections of elements + CREATE TABLE columns render STRUCT fields as Hive ``STRUCT``. + Compiling a type on its own renders ``ROW(...)``. + See Also: AWS Athena Data Types: https://docs.aws.amazon.com/athena/latest/ug/data-types.html @@ -119,7 +122,7 @@ def visit_TINYINT(self, type_: types.Integer, **kw: Any) -> str: return "TINYINT" def visit_INTEGER(self, type_: types.Integer, **kw: Any) -> str: - return "INT" if kw.get("_athena_array_ddl") else "INTEGER" + return "INT" if kw.get("_athena_hive_ddl") else "INTEGER" def visit_SMALLINT(self, type_: types.SmallInteger, **kw: Any) -> str: return "SMALLINT" @@ -197,31 +200,51 @@ 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 spelling for a CREATE TABLE column type. + + ``get_column_specification`` passes the column as ``type_expression``. + ARRAY compilation sets ``_athena_hive_ddl`` so nested fields use + ``STRUCT`` and ``INT``. STRUCT and MAP reuse that flag in + column DDL. Direct compilation and CAST leave it unset. + + Args: + kw: Type-compiler keyword arguments. When Hive spelling 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) - preparer = ( - AthenaDDLIdentifierPreparer(self.dialect) - if kw.get("_athena_array_ddl") - else self.dialect.identifier_preparer - ) - name = preparer.quote(field_name) - separator = ":" if kw.get("_athena_array_ddl") else " " - field_specs.append(f"{name}{separator}{field_type_str}") - if kw.get("_athena_array_ddl"): - return f"STRUCT<{', '.join(field_specs)}>" - return f"ROW({', '.join(field_specs)})" + # Empty structs keep the existing ROW() rendering in every context. + if not isinstance(type_, AthenaStruct) or not type_.fields: return "ROW()" - return "ROW()" + hive_ddl = self._enable_hive_column_ddl(kw) + preparer = ( + AthenaDDLIdentifierPreparer(self.dialect) + if hive_ddl + else self.dialect.identifier_preparer + ) + separator = ":" if hive_ddl else " " + field_specs = [] + for field_name, field_type in type_.fields.items(): + field_type_str = self.process(field_type, **kw) + field_specs.append(f"{preparer.quote(field_name)}{separator}{field_type_str}") + if hive_ddl: + return f"STRUCT<{', '.join(field_specs)}>" + return f"ROW({', '.join(field_specs)})" 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}>" @@ -232,7 +255,7 @@ def visit_MAP(self, type_, **kw): def visit_array(self, type_, **kw): if isinstance(type_, types.ARRAY): - kw["_athena_array_ddl"] = True + kw["_athena_hive_ddl"] = True item_type_str = self.process(_ArrayTypeInspector.item_type(type_), **kw) return f"ARRAY<{item_type_str}>" return "ARRAY" @@ -1088,6 +1111,7 @@ def get_column_specification(self, column: Column[Any], **kwargs) -> str: # use the int keyword to represent an integer type_ = "INT" else: + # type_expression marks column DDL so STRUCT and MAP use Hive syntax. type_ = self.dialect.type_compiler.process(column.type, type_expression=column) text = [f"{self.preparer.format_column(column)} {type_}"] if column.comment: diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index e85d183b..5dd7b1d5 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -2513,8 +2513,8 @@ 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 "metrics 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): @@ -2556,12 +2556,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.""" @@ -2599,6 +2599,57 @@ def test_create_table_with_complex_nested_types(self, engine): ) 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" + location = f"{ENV.s3_staging_dir}{ENV.schema}/{table_name}/" + 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=location, + awsathena_file_format="PARQUET", + ) + ddl = str(CreateTable(table).compile(dialect=conn.dialect)) + assert "profile STRUCT>" in ddl + assert "labels MAP>" in ddl + 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 df0e6b5b..db9f5baa 100644 --- a/tests/pyathena/sqlalchemy/test_compiler.py +++ b/tests/pyathena/sqlalchemy/test_compiler.py @@ -105,6 +105,31 @@ def test_visit_struct_single_field(self): result = compiler.visit_struct(struct_type) assert result == "ROW(name STRING)" or result == "ROW(name VARCHAR)" + def test_visit_struct_and_map_without_column_context_stay_row(self): + dialect = AthenaDialect() + compiler = AthenaTypeCompiler(dialect) + struct_type = AthenaStruct( + ("profile", AthenaStruct(("name", String), ("age", Integer))), + ("metrics", AthenaMap(String, Integer)), + ) + map_type = AthenaMap(Integer, AthenaStruct(("n", Integer))) + assert compiler.process(struct_type) == ( + "ROW(profile ROW(name STRING, age INTEGER), metrics MAP)" + ) + assert compiler.process(map_type) == "MAP" + assert compiler.process(AthenaStruct()) == "ROW()" + + def test_type_expression_column_selects_hive_syntax(self): + compiler = AthenaDialect().type_compiler_instance + struct_type = AthenaStruct(("name", String), ("age", Integer)) + map_type = AthenaMap(Integer, AthenaStruct(("n", Integer))) + assert compiler.process(struct_type, type_expression=Column("profile", struct_type)) == ( + "STRUCT" + ) + assert compiler.process(map_type, type_expression=Column("labels", map_type)) == ( + "MAP>" + ) + def test_visit_map_default(self): dialect = AthenaDialect() compiler = AthenaTypeCompiler(dialect) @@ -457,6 +482,43 @@ def test_stepped_array_aggregate_with_inferred_element_type(self): assert "Unsupported ARRAY slice step" in sql assert "ARRAY(NULL)" not in sql + @pytest.mark.parametrize( + ("type_", "expected"), + [ + ( + AthenaStruct(("name", String), ("age", Integer)), + "ROW(name VARCHAR, age INTEGER)", + ), + ( + AthenaStruct( + ("personal", AthenaStruct(("name", String), ("age", Integer))), + ("scores", types.ARRAY(Integer)), + ("attrs", AthenaMap(String, String)), + ), + "ROW(personal ROW(name VARCHAR, age INTEGER), scores ARRAY(INTEGER), " + "attrs MAP(VARCHAR, VARCHAR))", + ), + ( + AthenaMap(String, AthenaStruct(("value", String), ("count", Integer))), + "MAP(VARCHAR, ROW(value VARCHAR, count INTEGER))", + ), + ( + AthenaMap(String, AthenaMap(Integer, AthenaStruct(("n", Integer)))), + "MAP(VARCHAR, MAP(INTEGER, ROW(n INTEGER)))", + ), + ( + types.ARRAY(AthenaMap(String, AthenaStruct(("n", Integer)))), + "ARRAY(MAP(VARCHAR, ROW(n INTEGER)))", + ), + ( + AthenaStruct(("tags", AthenaMap(String, types.ARRAY(Integer)))), + "ROW(tags MAP(VARCHAR, ARRAY(INTEGER)))", + ), + ], + ) + def test_complex_cast_keeps_dml_syntax(self, type_, expected): + assert self._compile_sql(cast(column("col"), type_)) == f"CAST(col AS {expected})" + def _format_sql(self, statement, parameters=None): """Format a statement with the parameter names SQLAlchemy sends to the cursor. @@ -503,7 +565,10 @@ def test_like_pattern_is_not_truncated(self, pattern): 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 and nested MAP + types use 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 @@ -521,6 +586,16 @@ 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_s3tables_catalog_omits_location(self): table = Table( "tbl", @@ -632,3 +707,114 @@ def test_create_connect_args_stores_connect_options_for_subclass_dialects(self): ) ddl = str(CreateTable(table).compile(dialect=dialect)) assert "LOCATION" not in ddl + + def test_create_table_renders_hive_struct_syntax(self): + ddl = self._ddl( + Column("id", Integer), + Column("scores", AthenaArray(Integer)), + Column( + "profile", + AthenaStruct( + ("name", String), + ("age", Integer), + ("ratio", Float(24)), + ("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)), + ("scores", AthenaArray(Integer)), + ), + ), + Column( + "deep", + AthenaArray( + AthenaMap(String, AthenaStruct(("value", String), ("flag", types.Boolean))) + ), + ), + ) + assert "id INT" in ddl + assert "scores ARRAY" in ddl + assert ( + "profile STRUCT>" + ) in ddl + assert "labels MAP>" in ddl + assert "nested_maps MAP>>" in ddl + assert ( + "mixed STRUCT, attrs:MAP, scores:ARRAY>" + ) in ddl + assert "deep ARRAY>>" in ddl + assert "ROW(" not in ddl + assert "INTEGER" not in ddl + assert "REAL" not in ddl + assert "FLOAT(" not in ddl + + def test_struct_field_quoting_follows_ddl_preparer(self): + struct_type = AthenaStruct( + ("date", String), + ("select", Integer), + ('na"me', String), + ("a`b", String), + ("first name", String), + ("_hidden", Integer), + ) + ddl = self._ddl(Column("payload", struct_type)) + assert ( + 'payload STRUCT<`date`:STRING, `select`:INT, `na"me`:STRING, ' + "`a``b`:STRING, `first name`:STRING, `_hidden`:INT>" + ) in ddl + cast_sql = str(cast(column("payload"), struct_type).compile(dialect=AthenaDialect())) + assert cast_sql == ( + 'CAST(payload AS ROW(date VARCHAR, "select" INTEGER, "na""me" VARCHAR, ' + '"a`b" VARCHAR, "first name" VARCHAR, _hidden INTEGER))' + ) + + def test_empty_struct_column_stays_row(self): + ddl = self._ddl( + Column("empty", AthenaStruct()), + Column("filled", AthenaStruct(("n", Integer))), + ) + assert "empty ROW()" in ddl + assert "filled STRUCT" in ddl + assert "STRUCT<>" not in ddl + + def test_unsupported_type_inside_struct_column_still_raises(self): + with pytest.raises(exc.CompileError, match="not supported"): + self._ddl(Column("payload", AthenaStruct(("when", types.Time)))) + + def test_scalar_and_float_column_ddl_is_unchanged(self): + ddl = self._ddl( + Column("id", Integer), + Column("label", String), + Column("flag", types.Boolean), + Column("ratio", Float), + Column("ratio_prec", Float(24)), + Column("real_value", types.REAL), + Column("float_value", types.FLOAT), + Column("wide", types.Double), + Column("amount", Numeric(10, 2)), + ) + assert "id INT" in ddl + assert "label STRING" in ddl + assert "flag BOOLEAN" in ddl + assert "ratio FLOAT" in ddl + assert "ratio_prec FLOAT" in ddl + assert "real_value FLOAT" in ddl + assert "float_value FLOAT" in ddl + assert "wide DOUBLE" in ddl + assert "amount DECIMAL(10, 2)" in ddl + assert "INTEGER" not in ddl + assert "REAL" not in ddl + assert "FLOAT(" not in ddl From e3ec1a953c6eadadc8c498dd2e7956775a194b36 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Tue, 29 Sep 2026 22:17:20 +0900 Subject: [PATCH 20/20] Keep the ARRAY support working on SQLAlchemy 1.x (3.x only) 3.x declares sqlalchemy>=1.0.0, but the backported ARRAY support (#774) read two SQLAlchemy 2.0-only attributes. On SQLAlchemy 1.4.54, CAST to a STRUCT, MAP, or ARRAY type, ARRAY literals, and ARRAY element updates failed with AttributeError, although 3.x previously compiled those casts. - _complex_dml_type uses dialect.type_compiler instead of type_compiler_instance. SQLAlchemy 2.0 sets both to the same instance, and 3.x's get_column_specification already uses type_compiler. - A _variant_mapping() helper reads with_variant() variants from _variant_mapping on SQLAlchemy 2.0, or from the mapping of the Variant type decorator that SQLAlchemy 1.x returns. The eight direct _variant_mapping reads in array.py use it. On SQLAlchemy 2.0, the helper returns the same mapping, so behavior is unchanged. Master does not need this: it requires SQLAlchemy 2.0 since #858. Co-Authored-By: Claude Opus 5.5 --- pyathena/sqlalchemy/array.py | 36 +++++++++++++++++++------ pyathena/sqlalchemy/compiler.py | 2 +- tests/pyathena/sqlalchemy/test_array.py | 14 +++++++++- 3 files changed, 42 insertions(+), 10 deletions(-) diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index 09cefe69..358b6109 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -166,6 +166,25 @@ def __init__(self, element, type_): self.array_type = type_ +def _variant_mapping(type_: TypeEngine[Any]) -> Mapping[str, TypeEngine[Any]]: + """Return the dialect-specific variants of a type. + + SQLAlchemy 2.0 stores ``with_variant()`` variants in ``_variant_mapping``. + SQLAlchemy 1.x returns a ``Variant`` type decorator that keeps them in + ``mapping`` instead. + + Args: + type_: The type to inspect. + + Returns: + The variants keyed by dialect name, empty when the type has none. + """ + mapping = getattr(type_, "_variant_mapping", None) + if mapping is None: + mapping = getattr(type_, "mapping", None) + return mapping or {} + + class _ArrayTypeInspector: """Interpret nested ARRAY element types for SQL compilation and value conversion. @@ -196,8 +215,9 @@ def item_type(type_: sqltypes.ARRAY[Any]) -> TypeEngine[Any]: return type_.item_type def decorator_impl(self, type_: types.TypeDecorator[Any]) -> TypeEngine[Any]: - if self.dialect.name in type_._variant_mapping: - return type_._variant_mapping[self.dialect.name] + variant = _variant_mapping(type_).get(self.dialect.name) + if variant is not None: + return variant implementation = type_.load_dialect_impl(self.dialect) if isinstance(implementation, AthenaTimestamp): return types.TIMESTAMP() @@ -286,12 +306,12 @@ def _complex_values(value: Any, type_: TypeEngine[Any]): def _bind(self, value: Any, type_: TypeEngine[Any]) -> Any: if isinstance(type_, types.TypeDecorator): if ( - self.dialect.name not in type_._variant_mapping + self.dialect.name not in _variant_mapping(type_) and type(type_).bind_processor is not types.TypeDecorator.bind_processor ): processor = type_.bind_processor(self.dialect) return processor(value) if processor else value - if self.dialect.name not in type_._variant_mapping and type_._has_bind_processor: + if self.dialect.name not in _variant_mapping(type_) and type_._has_bind_processor: value = type_.process_bind_param(value, self.dialect) return self._bind(value, self._type_inspector.decorator_impl(type_)) if value is None: @@ -317,13 +337,13 @@ def _bind(self, value: Any, type_: TypeEngine[Any]) -> Any: def _literal(self, value: Any, type_: TypeEngine[Any]) -> str: if isinstance(type_, types.TypeDecorator): if ( - self.dialect.name not in type_._variant_mapping + self.dialect.name not in _variant_mapping(type_) and type(type_).literal_processor is not types.TypeDecorator.literal_processor ): literal_override = type_.literal_processor(self.dialect) if literal_override is not None: return literal_override(value) - if self.dialect.name not in type_._variant_mapping: + if self.dialect.name not in _variant_mapping(type_): if type_._has_literal_processor: value = type_.process_literal_param(value, self.dialect) elif type_._has_bind_processor: @@ -358,12 +378,12 @@ def _decode(self, value: Any, type_: TypeEngine[Any], as_tuple: bool = False) -> if isinstance(type_, types.TypeDecorator): value = self._decode(value, self._type_inspector.decorator_impl(type_), as_tuple) if ( - self.dialect.name not in type_._variant_mapping + self.dialect.name not in _variant_mapping(type_) and type(type_).result_processor is not types.TypeDecorator.result_processor ): processor = type_.result_processor(self.dialect, None) return processor(value) if processor else value - if self.dialect.name not in type_._variant_mapping and type_._has_result_processor: + if self.dialect.name not in _variant_mapping(type_) and type_._has_result_processor: return type_.process_result_value(value, self.dialect) return value if value is None: diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 590cf142..2a6c60c1 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -718,7 +718,7 @@ def _complex_dml_type(self, type_, *, require_precision=False): "ARRAY decimal values require explicit Numeric precision; " "specify precision and scale to avoid implicit rounding" ) - return self.dialect.type_compiler_instance.process(type_) + return self.dialect.type_compiler.process(type_) def visit_athena_array_json_projection(self, expression, **kw): value = self.process(expression.element, **kw) diff --git a/tests/pyathena/sqlalchemy/test_array.py b/tests/pyathena/sqlalchemy/test_array.py index 22a4624d..0ae07e3e 100644 --- a/tests/pyathena/sqlalchemy/test_array.py +++ b/tests/pyathena/sqlalchemy/test_array.py @@ -32,7 +32,7 @@ import pyathena from pyathena.formatter import DefaultParameterFormatter -from pyathena.sqlalchemy.array import _ArrayWriteIndexType +from pyathena.sqlalchemy.array import _ArrayWriteIndexType, _variant_mapping from pyathena.sqlalchemy.base import AthenaDialect from pyathena.sqlalchemy.types import ( ARRAY, @@ -299,6 +299,18 @@ def test_decorated_array_index_and_slice(self): assert processor('{"_pyathena_array":["1","2"]}') == (1, 2) +def test_variant_mapping(): + variant = _variant_mapping(types.String().with_variant(Integer(), "awsathena")) + assert list(variant) == ["awsathena"] + assert isinstance(variant["awsathena"], Integer) + assert _variant_mapping(types.String()) == {} + # SQLAlchemy 1.x types have no _variant_mapping; its Variant keeps them in mapping. + integer = Integer() + assert _variant_mapping(SimpleNamespace(mapping={"awsathena": integer})) == { + "awsathena": integer + } + + class TestArrayTypeInspector: @pytest.mark.parametrize( ("type_", "value", "expected"),