From fd80f984edd518cf3bfbaa98da970175ae061b02 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 12:59:47 +0900 Subject: [PATCH 1/3] Read managed query results as a CSV result file With managed query result storage, the pandas, Arrow, and Polars cursors built their results from GetQueryResults values converted by DefaultTypeConverter, so the cursor's converter, dtypes, and date parsing did not apply and the results differed from those of an S3 result file. GetQueryResults returns the same text for each value as the CSV result file, so the rows are now written in the format of that file and read with the same reader. The cursor's converter applies to them, including a custom one, and the types and values match the S3 path. Athena writes up to 12 fractional digits for timestamp columns, which the GetQueryResults fallback used to truncate to microseconds: - ArrowCursor reads timestamp columns as timestamp[us] instead of timestamp[ms], and always reads them as text and casts them, truncating the text when a value is longer than the unit holds. A 6-digit fraction failed to parse as timestamp[ms], and on Linux the "%Y-%m-%d %H:%M:%S %Z" strptime fallback accepted such values without their fraction. - PolarsCursor reads the timestamp columns again as text when a whole read fails to parse them, and truncates them to the time unit. Closes #1028 Closes #1042 Co-Authored-By: Claude Opus 5.5 --- docs/arrow.md | 2 +- docs/usage.md | 5 +- pyathena/arrow/converter.py | 2 +- pyathena/arrow/result_set.py | 116 +++++++++++------- pyathena/converter.py | 30 ----- pyathena/pandas/result_set.py | 118 +++++++++--------- pyathena/polars/result_set.py | 145 ++++++++++++++++------- pyathena/result_set.py | 119 +++++++++++-------- tests/pyathena/arrow/test_cursor.py | 35 +++++- tests/pyathena/arrow/test_result_set.py | 46 ++++++- tests/pyathena/pandas/test_cursor.py | 28 ++++- tests/pyathena/polars/test_cursor.py | 32 ++++- tests/pyathena/polars/test_result_set.py | 92 +++++++++++++- tests/pyathena/test_converter.py | 16 --- tests/pyathena/test_result_set.py | 24 ++++ tests/pyathena/util.py | 15 +++ 16 files changed, 571 insertions(+), 254 deletions(-) diff --git a/docs/arrow.md b/docs/arrow.md index 9cc66e731..5306f0902 100644 --- a/docs/arrow.md +++ b/docs/arrow.md @@ -148,7 +148,7 @@ class CustomArrowTypeConverter(Converter): "char": pa.string(), "varchar": pa.string(), "string": pa.string(), - "timestamp": pa.timestamp("ms"), + "timestamp": pa.timestamp("us"), "date": pa.timestamp("ms"), "time": pa.string(), "varbinary": pa.string(), diff --git a/docs/usage.md b/docs/usage.md index 492e7c11d..79d4b0422 100644 --- a/docs/usage.md +++ b/docs/usage.md @@ -42,6 +42,8 @@ cursor = connect(work_group="YOUR_MANAGED_WORK_GROUP", With managed query result storage, query results are retrieved via the `GetQueryResults` API (1000 rows per request) instead of reading S3 files directly. This may be slower for large result sets. For large datasets, consider using customer-managed storage or the `UNLOAD` statement. +The pandas, Arrow, and Polars cursors read these rows as they read a CSV result file, with the +same types and converters, but not in chunks. ``` ## Cursor iteration @@ -709,8 +711,7 @@ column. You can mix both styles in the same dictionary. JSON-formatted output, which is parsed reliably. - **Arrow, Pandas, and Polars cursors** — These cursors accept `result_set_type_hints` but their converters do not currently use the hints because they rely on their own - type systems. The parameter is passed through for forward compatibility and for - result sets that fall back to the default conversion path. + type systems. The parameter is passed through for forward compatibility. ### Breaking change in 3.30.0 diff --git a/pyathena/arrow/converter.py b/pyathena/arrow/converter.py index 7da7bcef3..51efec2fa 100644 --- a/pyathena/arrow/converter.py +++ b/pyathena/arrow/converter.py @@ -86,7 +86,7 @@ def _dtypes(self) -> dict[str, type[Any]]: "char": pa.string(), "varchar": pa.string(), "string": pa.string(), - "timestamp": pa.timestamp("ms"), + "timestamp": pa.timestamp("us"), "date": pa.timestamp("ms"), "time": pa.string(), "time with time zone": pa.string(), diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index d34a7cb50..97353d5a1 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -12,20 +12,45 @@ from pyathena import OperationalError from pyathena.arrow.util import to_column_info -from pyathena.converter import _TEXT_VALUE_TYPES, Converter, _text_value_converter, _to_default -from pyathena.error import ProgrammingError +from pyathena.converter import Converter, _to_default from pyathena.model import AthenaQueryExecution from pyathena.result_set import AthenaResultSet from pyathena.util import RetryConfig, override, parse_output_location if TYPE_CHECKING: import polars as pl - from pyarrow import Table + from pyarrow import ChunkedArray, Table, TimestampType from pyathena.connection import Connection _logger = logging.getLogger(__name__) +# The length of timestamp text, such as "2020-01-02 03:04:05.123456", that holds the +# fraction a timestamp unit can represent. +_TIMESTAMP_TEXT_LENGTHS: dict[str, int] = {"s": 19, "ms": 23, "us": 26, "ns": 29} + + +def _to_timestamp(column: ChunkedArray, type_: TimestampType) -> ChunkedArray: + """Convert timestamp text to a timestamp type, truncating finer fractions. + + Athena writes up to 12 fractional digits, which pyarrow does not parse into a + timestamp type whose unit holds fewer. + + Args: + column: The timestamp text, with NULL as null or as an empty string. + type_: The timestamp type. + + Returns: + The timestamps. + """ + import pyarrow as pa + import pyarrow.compute as pc + + length = _TIMESTAMP_TEXT_LENGTHS[type_.unit] + if (pc.max(pc.utf8_length(column)).as_py() or 0) > length: + column = pc.utf8_slice_codeunits(column, 0, length) + return pc.if_else(pc.equal(column, ""), pa.scalar(None, pa.string()), column).cast(type_) + class AthenaArrowResultSet(AthenaResultSet): """Result set that provides Apache Arrow Table results with columnar optimization. @@ -141,15 +166,13 @@ def __init__( if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: self._table = self._as_arrow() elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: - self._table = self._as_arrow_from_api() + # Without a result file, as with managed query result storage, the rows from + # GetQueryResults are read as a CSV result file. + self._table = self._read_csv() else: import pyarrow as pa self._table = pa.Table.from_pydict({}) - # The fetch methods convert the values read from a result file. GetQueryResults - # values are already converted, except json and time with time zone values, - # which stay text. - self._convert_rows = bool(self.output_location) self._batches = iter(self._table.to_batches(arraysize)) def _create_s3_file_system(self): @@ -258,12 +281,7 @@ def _fetch(self) -> None: # converters property keep one column per name. columns = [column.to_pylist() for column in rows.columns] description = self.description if self.description else [] - converters = [ - self._converter.get(d[1]) - if self._convert_rows or d[1] in _TEXT_VALUE_TYPES - else _to_default - for d in description - ] + converters = [self._converter.get(d[1]) for d in description] if any(convert is not _to_default for convert in converters): processed_rows = [ tuple(convert(v) for convert, v in zip(converters, row, strict=False)) @@ -287,12 +305,18 @@ def fetchone( return self._rows.popleft() def _read_csv(self) -> Table: + """Read the CSV result file, or the GetQueryResults rows as one without it. + + Returns: + The Arrow Table of the results. + + Raises: + OperationalError: If reading the results fails. + """ import pyarrow as pa from pyarrow import csv - if not self.output_location: - raise ProgrammingError("OutputLocation is none or empty.") - if not self.output_location.endswith((".csv", ".txt")): + if self.output_location and not self.output_location.endswith((".csv", ".txt")): return pa.Table.from_pydict({}) if self.substatement_type and self.substatement_type.upper() in ( "UPDATE", @@ -301,7 +325,14 @@ def _read_csv(self) -> Table: "VACUUM_TABLE", ): return pa.Table.from_pydict({}) - length = self._get_content_length() + if self.output_location: + data = None + length = self._get_content_length() + location = "/".join(parse_output_location(self.output_location)) + else: + data = self._fetch_all_rows_as_csv() + length = len(data) + location = "the GetQueryResults rows" description = self.description if self.description else [] names = [d[0] for d in description] # pyarrow types every column with a name by its column_types entry, so columns @@ -318,8 +349,16 @@ def _read_csv(self) -> Table: else: column_names = names column_types = self.column_types + # Timestamp columns are read as text: pyarrow does not parse fractions finer + # than the unit, and on some platforms its strptime fallbacks accept such a + # value without its fraction. + timestamp_types = { + i: dtype + for i, (name, d) in enumerate(zip(column_names, description, strict=True)) + if d[1] == "timestamp" and isinstance(dtype := column_types.get(name), pa.TimestampType) + } binary_columns = {i for i, d in enumerate(description) if d[1] == "varbinary"} - if length and self.output_location.endswith(".txt"): + if length and self.output_location and self.output_location.endswith(".txt"): read_opts = csv.ReadOptions( skip_rows=0, column_names=column_names, @@ -332,7 +371,7 @@ def _read_csv(self) -> Table: double_quote=False, escape_char=False, ) - elif length and self.output_location.endswith(".csv"): + elif length: read_opts = csv.ReadOptions(skip_rows=0, block_size=self._block_size, use_threads=True) if has_duplicate_names: read_opts.column_names = column_names @@ -353,19 +392,27 @@ def _read_csv(self) -> Table: else: return pa.Table.from_pydict({}) - bucket, key = parse_output_location(self.output_location) try: table = csv.read_csv( - self._fs.open_input_stream(f"{bucket}/{key}"), + self._fs.open_input_stream(location) if data is None else pa.BufferReader(data), 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=column_types, + column_types={ + **column_types, + **{column_names[i]: pa.string() for i in timestamp_types}, + }, ), ) + for index, type_ in timestamp_types.items(): + table = table.set_column( + index, + pa.field(table.schema.field(index).name, type_), + _to_timestamp(table.column(index), type_), + ) if has_duplicate_names: table = table.rename_columns(names) if binary_columns: @@ -377,7 +424,7 @@ def _read_csv(self) -> Table: 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}.") + _logger.exception(f"Failed to read {location}.") raise OperationalError(*e.args) from e def _read_parquet(self) -> Table: @@ -406,27 +453,6 @@ def _as_arrow(self) -> Table: table = self._read_csv() return table - def _as_arrow_from_api(self, converter: Converter | None = None) -> Table: - """Build an Arrow Table from GetQueryResults API. - - Used as a fallback when ``output_location`` is not available - (e.g. managed query result storage). - - Args: - converter: Type converter for result values. Defaults to - ``DefaultTypeConverter`` with json and time with time zone values kept as - text, as in the CSV result file. Arrow has no type for JSON values or for - times with a time zone. - """ - import pyarrow as pa - - rows = self._fetch_all_rows(converter or _text_value_converter()) - if not rows: - return pa.Table.from_pydict({}) - description = self.description if self.description else [] - columns = [list(column) for column in zip(*rows, strict=True)] - return pa.table(columns, names=[d[0] for d in description]) - def as_arrow(self) -> Table: """Return the query results as an Apache Arrow Table. diff --git a/pyathena/converter.py b/pyathena/converter.py index 72d9d83de..a12513c37 100644 --- a/pyathena/converter.py +++ b/pyathena/converter.py @@ -803,33 +803,3 @@ def _parse_type_hint(self, type_hint: str) -> TypeNode: if normalized not in self._parsed_hints: self._parsed_hints[normalized] = self._parser.parse(normalized) return self._parsed_hints[normalized] - - -# The types whose values the Arrow and Polars GetQueryResults fallbacks keep as text, -# as in a CSV result file, and convert when the rows are fetched. -_TEXT_VALUE_TYPES: tuple[str, ...] = ("json", "time with time zone", "timestamp with time zone") - - -def _text_value_converter() -> DefaultTypeConverter: - """Return a ``DefaultTypeConverter`` that keeps ``_TEXT_VALUE_TYPES`` values as text. - - Values nested in typed complex values keep only the time zone types as text, - because Arrow and Polars time and timestamp types hold one time zone per column; - nested JSON values are decoded as before. - - Returns: - The converter. - """ - converter = DefaultTypeConverter() - for type_ in _TEXT_VALUE_TYPES: - converter.set(type_, _to_default) - converter._typed_converter = TypedValueConverter( - converters={ - **_DEFAULT_CONVERTERS, - "time with time zone": _to_default, - "timestamp with time zone": _to_default, - }, - default_converter=_to_default, - struct_parser=_to_struct, - ) - return converter diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 2ae910785..8e68adc6b 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -8,7 +8,7 @@ from collections.abc import Callable, Iterable, Iterator from contextlib import ExitStack from functools import partial -from io import BufferedReader, IOBase, StringIO, TextIOWrapper +from io import BufferedReader, BytesIO, IOBase, StringIO, TextIOWrapper from multiprocessing import cpu_count from typing import ( TYPE_CHECKING, @@ -421,7 +421,6 @@ class AthenaPandasResultSet(AthenaResultSet): AUTO_CHUNK_SIZE_LARGE: int = 100_000 AUTO_CHUNK_SIZE_MEDIUM: int = 50_000 - _INTEGER_TYPES: ClassVar[tuple[str, ...]] = ("tinyint", "smallint", "integer", "bigint") _PARSE_DATES: ClassVar[list[str]] = [ "date", "time", @@ -546,16 +545,19 @@ def __init__( # The whole result when it was not read in chunks. self._df: DataFrame | None = None - if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: - result = self._as_pandas() - trunc_date = _no_trunc_date if self.is_unload else self._finish_csv_frame + if self.state == AthenaQueryExecution.STATE_SUCCEEDED: + if self.output_location: + result = self._as_pandas() + trunc_date = _no_trunc_date if self.is_unload else self._finish_csv_frame + else: + # Without a result file, as with managed query result storage, the rows + # from GetQueryResults are read as a CSV result file, but not in chunks. + result = self._read_csv() + trunc_date = self._finish_csv_frame if isinstance(result, pd.DataFrame): self._df = trunc_date(result) else: self._df_iter = PandasDataFrameIterator(result, trunc_date, self._csv_stream) - elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: - # GetQueryResults values are already converted and need no time truncation. - self._df = self._as_pandas_from_api() else: self._df = pd.DataFrame() if self._df is not None: @@ -795,11 +797,19 @@ def fetchone( return tuple(row[1].values()) def _read_csv(self) -> TextFileReader | DataFrame: + """Read the CSV result file, or the GetQueryResults rows as one without it. + + The GetQueryResults rows are read as a whole, not in chunks. + + Returns: + The DataFrame of the results, or a reader of their chunks. + + Raises: + OperationalError: If reading the results fails. + """ import pandas as pd - if not self.output_location: - raise ProgrammingError("OutputLocation is none or empty.") - if not self.output_location.endswith((".csv", ".txt")): + if self.output_location and not self.output_location.endswith((".csv", ".txt")): return pd.DataFrame() if self.substatement_type and self.substatement_type.upper() in ( "UPDATE", @@ -808,15 +818,22 @@ def _read_csv(self) -> TextFileReader | DataFrame: "VACUUM_TABLE", ): return pd.DataFrame() - length = self._get_content_length() + if self.output_location: + data = None + length = self._get_content_length() + location = self.output_location + else: + data = self._fetch_all_rows_as_csv() + length = len(data) + location = "the GetQueryResults rows" if length == 0: return pd.DataFrame() # Chunksize determination with user preference priority - effective_chunksize = self._chunksize + effective_chunksize = self._chunksize if data is None else None # Only auto-optimize if user hasn't specified chunksize AND auto_optimize is enabled - if effective_chunksize is None and self._auto_optimize_chunksize: + if effective_chunksize is None and self._auto_optimize_chunksize and data is None: effective_chunksize = self._auto_determine_chunksize(length) if effective_chunksize: _logger.debug( @@ -832,13 +849,22 @@ def _read_csv(self) -> TextFileReader | DataFrame: try: with ExitStack() as stack: - source: str | IOBase = self.output_location + source: str | IOBase = location binary_columns = self._configure_binary_csv_read(read_csv_kwargs, labels) if labels is not None: # After _configure_binary_csv_read(), which checks for the header row. self._read_csv_header_as_labels(read_csv_kwargs, csv_engine) self._csv_converters = read_csv_kwargs.get("converters") or {} - if binary_columns: + if data is not None: + # The rows are in memory, with nothing to open with storage options. + read_csv_kwargs.pop("storage_options", None) + if binary_columns: + source = self._csv_stream = stack.enter_context( + self._open_binary_csv_stream(binary_columns, None, data) + ) + else: + source = BytesIO(data) + elif binary_columns: # Given storage_options, even None, open the file through fsspec # as pandas does. storage_options = None @@ -861,7 +887,7 @@ def _read_csv(self) -> TextFileReader | DataFrame: stack.pop_all() # Log performance information for large files - if length > self.LARGE_FILE_THRESHOLD_BYTES: + if data is None and 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( @@ -872,7 +898,7 @@ def _read_csv(self) -> TextFileReader | DataFrame: return result except Exception as e: - _logger.exception(f"Failed to read {self.output_location}.") + _logger.exception(f"Failed to read {location}.") raise OperationalError(*e.args) from e def _reads_csv_with_pyarrow(self) -> bool: @@ -984,8 +1010,7 @@ def _is_standard_csv_parsing(self, read_csv_kwargs: dict[str, Any]) -> bool: True if pandas reads the header and quoted fields of the file as written. """ return not ( - not self.output_location - or not self.output_location.endswith(".csv") + (self.output_location and 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 @@ -1150,12 +1175,28 @@ def _configure_binary_csv_read( return binary_columns def _open_binary_csv_stream( - self, binary_columns: set[int], storage_options: dict[str, Any] | None + self, + binary_columns: set[int], + storage_options: dict[str, Any] | None, + data: bytes | None = None, ) -> TextIOWrapper: - """Open a stream that preserves binary NULL fields and original CSV newlines.""" + """Open a stream that preserves binary NULL fields and original CSV newlines. + + Args: + binary_columns: Zero-based indexes of the binary columns. + storage_options: The options to open the result file with through fsspec, + or None to open it with PyAthena's filesystem. + data: The CSV text to read in place of the result file. + + Returns: + The text stream for ``pandas.read_csv()``. + """ text_options: dict[str, Any] = {"mode": "rt", "encoding": "utf-8", "newline": ""} with ExitStack() as stack: - if storage_options is None: + source: IOBase + if data is not None: + source = StringIO(data.decode("utf-8"), newline="") + elif storage_options is None: source = stack.enter_context(self._fs.open(self.output_location, **text_options)) else: source = stack.enter_context( @@ -1227,37 +1268,6 @@ def _as_pandas(self) -> TextFileReader | DataFrame: df = self._read_csv() return df - def _as_pandas_from_api(self, converter: Converter | None = None) -> DataFrame: - """Build a DataFrame from GetQueryResults API. - - Used as a fallback when ``output_location`` is not available - (e.g. managed query result storage). - - Args: - converter: Type converter for result values. Defaults to - ``DefaultTypeConverter`` if not specified. - """ - import pandas as pd - - rows = self._fetch_all_rows(converter) - if not rows: - return pd.DataFrame() - description = self.description if self.description else [] - # Positional, so that columns with the same name keep their own values. - columns = [list(column) for column in zip(*rows, strict=True)] - # Integer columns get the dtype that the CSV result file reads them with, - # and json columns with NULL stay objects as there, so that NULL does not - # make their values floats. - data: dict[Any, Any] = {} - for name, values, d in zip(self._get_column_names(), columns, description, strict=True): - dtype = None - if d[1] in self._INTEGER_TYPES: - dtype = self._converter.get_dtype(d[1], d[4], d[5]) - elif d[1] == "json" and None in values: - dtype = object - data[name] = values if dtype is None else pd.array(values, dtype=dtype) - return pd.DataFrame(data) - def as_pandas(self) -> PandasDataFrameIterator | DataFrame: """Return the query results as a DataFrame or an iterator of DataFrame chunks. diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index 85ea4742b..4ef0c750f 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -22,7 +22,7 @@ ) from pyathena import OperationalError -from pyathena.converter import Converter, _text_value_converter +from pyathena.converter import Converter from pyathena.error import ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.polars.util import to_column_info @@ -38,11 +38,48 @@ _logger = logging.getLogger(__name__) +# The length of timestamp text, such as "2020-01-02 03:04:05.123456", that holds the +# fraction a Datetime time unit can represent. +_TIMESTAMP_TEXT_LENGTHS: dict[str, int] = {"ms": 23, "us": 26, "ns": 29} + + def _identity(x: Any) -> Any: """Identity function for use as default converter.""" return x +def _to_datetimes(df: pl.DataFrame, dtypes: dict[str, Any]) -> pl.DataFrame: + """Convert timestamp text columns to Datetime dtypes, truncating finer fractions. + + Athena writes up to 12 fractional digits, which Polars does not parse into a + Datetime whose time unit holds fewer. + + Args: + df: The DataFrame with the timestamp text columns, with NULL as null or as + an empty string. + dtypes: The Datetime dtypes keyed by column name. + + Returns: + The DataFrame with the columns converted. + """ + import polars as pl + + exprs = [] + for name, dtype in dtypes.items(): + if name not in df.columns: + # Not selected, as with ``columns`` given to ``execute()``. + continue + time_unit = (dtype() if isinstance(dtype, type) else dtype).time_unit + text = pl.col(name) + exprs.append( + pl.when(text != "") + .then(text.str.slice(0, _TIMESTAMP_TEXT_LENGTHS[time_unit])) + .str.to_datetime("%Y-%m-%d %H:%M:%S%.f", time_unit=time_unit) + .cast(dtype) + ) + return df.with_columns(exprs) if exprs else df + + class PolarsDataFrameIterator(abc.Iterator): # type: ignore[type-arg] """Iterator for chunked DataFrame results from Athena queries. @@ -278,13 +315,10 @@ def __init__( else: self._df_iter = self._create_dataframe_iterator() elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: - self._df = self._as_polars_from_api() - # GetQueryResults values are already converted, except json and time with - # time zone values kept as text. - column_names = self._get_frame_column_names() - self._df_converters = self._text_value_converters( - self._get_converters(column_names), column_names - ) + # Without a result file, as with managed query result storage, the rows from + # GetQueryResults are read as a CSV result file, but not in chunks. + self._df = self._read_csv() + self._df_converters = self._get_converters(self._get_frame_column_names()) else: self._df = pl.DataFrame() if self._df is not None: @@ -432,6 +466,30 @@ def _get_converters( for name, d in zip(column_names, description, strict=True) } + def _get_timestamp_dtypes(self, has_header: bool) -> dict[str, Any]: + """Get the Datetime dtypes of the timestamp columns, which are read as text. + + Args: + has_header: Whether the CSV data has a header. Without one, or with + ``schema_overrides`` or ``with_column_names`` given to ``execute()``, + Polars does not name the columns by the header, and they are not read + as text. + + Returns: + The Datetime dtypes keyed by the header of a CSV file. + """ + import polars as pl + + if not has_header or self._kwargs.keys() & {"schema_overrides", "with_column_names"}: + return {} + dtypes = self._csv_dtypes + return { + name: dtype + for name, d in zip(self._get_column_names(), self.description or [], strict=True) + if d[1] == "timestamp" + and ((dtype := dtypes.get(name)) is pl.Datetime or isinstance(dtype, pl.Datetime)) + } + def _get_column_names(self) -> list[str]: """Get the names of the result columns in a DataFrame. @@ -462,7 +520,7 @@ def _get_frame_column_names(self) -> list[str]: """ names = self._get_column_names() new_columns = self._kwargs.get("new_columns") - if not new_columns or not self.output_location or self.is_unload: + if not new_columns or (self.output_location and self.is_unload): return names return [*new_columns[: len(names)], *names[len(new_columns) :]] @@ -507,15 +565,12 @@ def fetchone( def _is_csv_readable(self) -> bool: """Check if CSV output is available and can be read. + Without an output location, the GetQueryResults rows are read as CSV. + Returns: True if CSV data is available to read, False otherwise. - - Raises: - ProgrammingError: If output location is not set. """ - if not self.output_location: - raise ProgrammingError("OutputLocation is none or empty.") - if not self.output_location.endswith((".csv", ".txt")): + if self.output_location and not self.output_location.endswith((".csv", ".txt")): return False if self.substatement_type and self.substatement_type.upper() in ( "UPDATE", @@ -524,8 +579,7 @@ def _is_csv_readable(self) -> bool: "VACUUM_TABLE", ): return False - length = self._get_content_length() - return length != 0 + return not self.output_location or self._get_content_length() != 0 def _prepare_parquet_location(self) -> bool: """Prepare unload location for Parquet reading. @@ -541,25 +595,21 @@ def _prepare_parquet_location(self) -> bool: return True def _read_csv(self) -> pl.DataFrame: - """Read query results from CSV file in S3. + """Read query results from CSV file in S3, or the GetQueryResults rows as one without it. Returns: Polars DataFrame containing the CSV data. Raises: - ProgrammingError: If output location is not set. - OperationalError: If reading the CSV file fails. + OperationalError: If reading the CSV data fails. """ import polars as pl if not self._is_csv_readable(): return pl.DataFrame() - if self.output_location is None: - raise ProgrammingError("output_location is not available.") - separator, has_header, new_columns = self._get_csv_params() - read_kwargs = self._read_kwargs( + kwargs = self._read_kwargs( lambda: self._csv_storage_options, separator=separator, has_header=has_header, @@ -567,15 +617,37 @@ def _read_csv(self) -> pl.DataFrame: ) if "schema_overrides" not in self._kwargs: # Renamed after reading, so that Polars matches the types to the header. - read_kwargs.pop("new_columns", None) + kwargs.pop("new_columns", None) + source: str | bytes + if self.output_location: + source = location = self.output_location + else: + source = self._fetch_all_rows_as_csv() + if not source: + return pl.DataFrame() + location = "the GetQueryResults rows" + del kwargs["storage_options"] try: - df = pl.read_csv(self.output_location, **read_kwargs) + try: + df = pl.read_csv(source, **kwargs) + except pl.exceptions.ComputeError: + timestamp_dtypes = self._get_timestamp_dtypes(has_header) + if not timestamp_dtypes: + raise + # Athena writes up to 12 fractional digits, which Polars does not parse + # into a Datetime whose time unit holds fewer, so the data is read again + # with the timestamp columns as text. + kwargs["schema_overrides"] = { + **self._csv_dtypes, + **dict.fromkeys(timestamp_dtypes, pl.String), + } + df = _to_datetimes(pl.read_csv(source, **kwargs), timestamp_dtypes) if new_columns: df.columns = [*new_columns, *df.columns[len(new_columns) :]] return df except Exception as e: - _logger.exception(f"Failed to read {self.output_location}.") + _logger.exception(f"Failed to read {location}.") raise OperationalError(*e.args) from e def _read_parquet(self) -> pl.DataFrame: @@ -641,25 +713,6 @@ def _as_polars(self) -> pl.DataFrame: df = self._read_csv() return df - def _as_polars_from_api(self, converter: Converter | None = None) -> pl.DataFrame: - """Build a Polars DataFrame from GetQueryResults API. - - Used as a fallback when ``output_location`` is not available - (e.g. managed query result storage). - - Args: - converter: Type converter for result values. Defaults to - ``DefaultTypeConverter`` with json and time with time zone values kept as - text, as in the CSV result file. A Polars ``Time`` has no time zone. - """ - import polars as pl - - rows = self._fetch_all_rows(converter or _text_value_converter()) - if not rows: - return pl.DataFrame() - columns = [list(column) for column in zip(*rows, strict=True)] - return pl.DataFrame(dict(zip(self._get_column_names(), columns, strict=True))) - def as_polars(self) -> pl.DataFrame: """Return query results as a Polars DataFrame. diff --git a/pyathena/result_set.py b/pyathena/result_set.py index c09117782..fd97ac384 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -12,13 +12,13 @@ ) from pyathena.common import BaseCursor, CursorIterator -from pyathena.converter import _TEXT_VALUE_TYPES, Converter, DefaultTypeConverter +from pyathena.converter import Converter, DefaultTypeConverter from pyathena.error import DataError, OperationalError, ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.util import RetryConfig, override, parse_output_location, retry_api_call if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Iterator from pyathena.connection import Connection @@ -675,29 +675,34 @@ def _is_first_row_column_labels(self, rows: list[dict[str, Any]]) -> bool: return False return True - def _text_value_converters( - self, - converters: dict[str, Callable[[str | None], Any | None]], - column_names: list[str] | None = None, - ) -> dict[str, Callable[[str | None], Any | None]]: - """Select the converters of the columns that the fallbacks keep as text. + def _iter_all_row_pages(self) -> Iterator[tuple[list[dict[str, Any]], int]]: + """Fetch all rows via GetQueryResults API from the beginning. - Args: - converters: The converters keyed by column name. - column_names: The names that ``converters`` uses for the columns, in - column order. Defaults to the names in the description. + Paginates through all results using MaxResults=1000. This is for subclass result + sets that need to fall back to the API when S3 output is not available (e.g., + managed query result storage). - Returns: - The converters of the columns whose Athena type is in ``_TEXT_VALUE_TYPES``. + Yields: + The rows of each page and the offset of their first data row, which is 1 when + the first page starts with the column labels. """ - description = self.description if self.description else [] - if column_names is None: - column_names = [d[0] for d in description] - return { - name: converters[name] - for name, d in zip(column_names, description, strict=True) - if d[1] in _TEXT_VALUE_TYPES - } + _logger.warning( + "output_location is not available (e.g. managed query result storage). " + "Falling back to GetQueryResults API. " + "This may be slow for large result sets." + ) + + next_token: str | None = None + first_page = True + while True: + response = self._get_query_results(self.DEFAULT_FETCH_SIZE, next_token) + rows, next_token = self._parse_result_rows(response) + # Only the first page starts with the column labels. + offset = 1 if first_page and rows and self._is_first_row_column_labels(rows) else 0 + first_page = False + yield rows, offset + if not next_token: + break def _fetch_all_rows( self, @@ -705,53 +710,63 @@ def _fetch_all_rows( ) -> list[tuple[Any | None, ...]]: """Fetch all rows via GetQueryResults API with type conversion. - Paginates through all results from the beginning using MaxResults=1000. - Defaults to ``DefaultTypeConverter`` for string-to-Python type conversion, - because subclass converters (e.g. Pandas/Arrow) are designed for S3 file - reading and may not handle API result strings. - - This method is intended for use by subclass result sets that need to - fall back to the API when S3 output is not available (e.g., managed - query result storage). - Args: converter: Type converter for result values. Defaults to ``DefaultTypeConverter`` if not specified. Returns: List of converted row tuples. + + Raises: + ProgrammingError: If the metadata is not available. """ - if self._metadata is None: + metadata = self._metadata + if metadata is None: raise ProgrammingError("Metadata is not available.") - - _logger.warning( - "output_location is not available (e.g. managed query result storage). " - "Falling back to GetQueryResults API. " - "This may be slow for large result sets." - ) - converter = converter or DefaultTypeConverter() all_rows: list[tuple[Any | None, ...]] = [] - next_token: str | None = None - - while True: - first_page = next_token is None - response = self._get_query_results(self.DEFAULT_FETCH_SIZE, next_token) - rows, next_token = self._parse_result_rows(response) - - # Only the first page can start with the column labels. - offset = 1 if first_page and rows and self._is_first_row_column_labels(rows) else 0 + for rows, offset in self._iter_all_row_pages(): all_rows.extend( cast( list[tuple[Any | None, ...]], - self._get_rows(offset, self._metadata, rows, converter), + self._get_rows(offset, metadata, rows, converter), ) ) + return all_rows - if not next_token: - break + def _fetch_all_rows_as_csv(self) -> bytes: + """Fetch all rows via GetQueryResults API as the text of a CSV result file. - return all_rows + GetQueryResults returns the same text for each value as the CSV result file + has, and this writes it in the same format: a header of the column labels, + each value quoted with its quotes doubled, NULL as an empty field, and a line + feed after each row. The cursors that read CSV result files can then read it + as one. + + Returns: + The CSV text encoded in UTF-8, or empty bytes when the result has no columns. + + Raises: + ProgrammingError: If the metadata is not available. + """ + if self._metadata is None: + raise ProgrammingError("Metadata is not available.") + description = self.description if self.description else [] + if not description: + return b"" + lines = [",".join('"' + d[0].replace('"', '""') + '"' for d in description)] + for rows, offset in self._iter_all_row_pages(): + lines.extend( + ",".join( + "" + if (value := data.get("VarCharValue")) is None + else '"' + value.replace('"', '""') + '"' + for data in row.get("Data", []) + ) + for row in rows[offset:] + ) + lines.append("") + return "\n".join(lines).encode("utf-8") def _get_content_length(self) -> int: if not self.output_location: diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index 54bc92536..54352fa9c 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -27,7 +27,11 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.util import CONVERTED_VALUES_QUERY, CONVERTED_VALUES_ROW +from tests.pyathena.util import ( + CONVERTED_VALUES_QUERY, + CONVERTED_VALUES_ROW, + RESULT_FILE_VALUES_QUERY, +) class TestArrowCursor: @@ -373,7 +377,7 @@ def test_complex_as_arrow(self, arrow_cursor): pa.field("col_double", pa.float64()), pa.field("col_string", pa.string()), pa.field("col_varchar", pa.string()), - pa.field("col_timestamp", pa.timestamp("ms")), + pa.field("col_timestamp", pa.timestamp("us")), pa.field("col_time", pa.string()), pa.field("col_date", pa.timestamp("ms")), pa.field("col_binary", pa.string()), @@ -550,7 +554,7 @@ def test_complex_as_polars(self, arrow_cursor): pl.Float64, pl.String, pl.String, - pl.Datetime("ms"), + pl.Datetime("us"), pl.String, pl.Datetime("ms"), pl.String, @@ -1055,6 +1059,31 @@ def test_fetch_all_rows(self, arrow_cursor): "2024-02-29 23:59:58.123 +05:30" ] + @pytest.mark.skipif(not ENV.managed_work_group, reason="AWS_ATHENA_MANAGED_WORKGROUP not set") + def test_managed_results_match_result_file(self): + """Managed results have the types and values of the CSV result file. + + The cursor's converter applies to them, including a custom one. + """ + results = [] + for kwargs in ({}, {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}): + converter = DefaultArrowTypeConverter() + converter.set("varchar", lambda value: value.upper() if value else value) + with ( + contextlib.closing(connect(**kwargs)) as conn, + conn.cursor(ArrowCursor, converter=converter) as cursor, + ): + cursor.execute(RESULT_FILE_VALUES_QUERY) + results.append((cursor.as_arrow(), cursor.fetchall())) + (table, rows), (managed_table, managed_rows) = results + assert managed_table.schema == table.schema + assert managed_table.equals(table) + assert managed_rows == rows + assert rows[0][1] == 'A,"B"\nC' + assert table.schema.field("col_timestamp_6").type == pa.timestamp("us") + # Fractions finer than microseconds are truncated. + assert (rows[0][5], rows[0][13]) == (datetime(2020, 1, 2, 3, 4, 5, 123456),) * 2 + @pytest.mark.parametrize( "arrow_cursor", [ diff --git a/tests/pyathena/arrow/test_result_set.py b/tests/pyathena/arrow/test_result_set.py index c79aa1e29..62f571486 100644 --- a/tests/pyathena/arrow/test_result_set.py +++ b/tests/pyathena/arrow/test_result_set.py @@ -4,14 +4,58 @@ # See LICENSE or https://opensource.org/licenses/MIT. # # SPDX-License-Identifier: MIT +from datetime import datetime from unittest.mock import MagicMock, patch +import pyarrow as pa +import pytest + from pyathena.arrow.converter import DefaultArrowTypeConverter -from pyathena.arrow.result_set import AthenaArrowResultSet +from pyathena.arrow.result_set import AthenaArrowResultSet, _to_timestamp from pyathena.model import AthenaQueryExecution from pyathena.util import RetryConfig +@pytest.mark.parametrize( + ("unit", "microseconds"), + [ + ("s", [0, 0, 0, 0, 0]), + ("ms", [0, 123000, 123000, 123000, 123000]), + ("us", [0, 123000, 123456, 123456, 123456]), + ], +) +def test_to_timestamp(unit, microseconds): + """Timestamp text with up to 12 fractional digits is truncated to the unit. + + NULL can be null or an empty string, depending on the CSV read options. + """ + column = pa.chunked_array( + [ + pa.array( + [ + "2020-01-02 03:04:05", + "2020-01-02 03:04:05.123", + "2020-01-02 03:04:05.123456", + "2020-01-02 03:04:05.123456789", + "2020-01-02 03:04:05.123456789012", + "0001-01-01 00:00:00.000", + "", + None, + ], + pa.string(), + ) + ] + ) + values = _to_timestamp(column, pa.timestamp(unit)) + assert values.type == pa.timestamp(unit) + assert values.to_pylist() == [ + *(datetime(2020, 1, 2, 3, 4, 5, us) for us in microseconds), + datetime(1, 1, 1), + None, + None, + ] + + class TestAthenaArrowResultSet: def test_fetch_after_close(self): """No AWS calls; the query execution and the filesystem are mocked.""" diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 44bbf2598..311c0bed8 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -27,7 +27,12 @@ from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.util import CONVERTED_VALUES_QUERY, CONVERTED_VALUES_ROW, cached_file_systems +from tests.pyathena.util import ( + CONVERTED_VALUES_QUERY, + CONVERTED_VALUES_ROW, + RESULT_FILE_VALUES_QUERY, + cached_file_systems, +) def _pandas_converter_without_bigint_dtype(): @@ -1750,6 +1755,27 @@ def test_fetch_all_rows(self, pandas_cursor): pandas_cursor.execute(CONVERTED_VALUES_QUERY) assert pandas_cursor.fetchall() == [CONVERTED_VALUES_ROW] + @pytest.mark.skipif(not ENV.managed_work_group, reason="AWS_ATHENA_MANAGED_WORKGROUP not set") + def test_managed_results_match_result_file(self): + """Managed results have the dtypes and values of the CSV result file. + + The cursor's converter applies to them, including a custom one. + """ + results = [] + for kwargs in ({}, {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}): + converter = DefaultPandasTypeConverter() + converter.set("varchar", lambda value: value.upper() if value else value) + with ( + contextlib.closing(connect(**kwargs)) as conn, + conn.cursor(PandasCursor, converter=converter) as cursor, + ): + cursor.execute(RESULT_FILE_VALUES_QUERY) + results.append((cursor.as_pandas(), cursor.fetchone())) + (df, row), (managed_df, managed_row) = results + pd.testing.assert_frame_equal(managed_df, df) + assert managed_row == row + assert row[1] == 'A,"B"\nC' + @pytest.mark.parametrize( "pandas_cursor", [ diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index d4c642eb6..d739bc22f 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -19,12 +19,18 @@ from pyathena.error import DatabaseError, ProgrammingError from pyathena.model import AthenaQueryExecution +from pyathena.polars.converter import DefaultPolarsTypeConverter from pyathena.polars.cursor import PolarsCursor from pyathena.polars.result_set import AthenaPolarsResultSet from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect -from tests.pyathena.util import CONVERTED_VALUES_QUERY, CONVERTED_VALUES_ROW, cached_file_systems +from tests.pyathena.util import ( + CONVERTED_VALUES_QUERY, + CONVERTED_VALUES_ROW, + RESULT_FILE_VALUES_QUERY, + cached_file_systems, +) class TestPolarsCursor: @@ -787,6 +793,30 @@ def test_fetch_all_rows(self, polars_cursor): "2024-02-29 23:59:58.123 +05:30" ] + @pytest.mark.skipif(not ENV.managed_work_group, reason="AWS_ATHENA_MANAGED_WORKGROUP not set") + def test_managed_results_match_result_file(self): + """Managed results have the types and values of the CSV result file. + + The cursor's converter applies to them, including a custom one. + """ + results = [] + for kwargs in ({}, {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}): + converter = DefaultPolarsTypeConverter() + converter.set("varchar", lambda value: value.upper() if value else value) + with ( + contextlib.closing(connect(**kwargs)) as conn, + conn.cursor(PolarsCursor, converter=converter) as cursor, + ): + cursor.execute(RESULT_FILE_VALUES_QUERY) + results.append((cursor.as_polars(), cursor.fetchall())) + (df, rows), (managed_df, managed_rows) = results + assert managed_df.schema == df.schema + assert managed_df.equals(df) + assert managed_rows == rows + assert rows[0][1] == 'A,"B"\nC' + # Fractions finer than microseconds are truncated. + assert (rows[0][5], rows[0][13]) == (datetime(2020, 1, 2, 3, 4, 5, 123456),) * 2 + @pytest.mark.parametrize( "polars_cursor", [ diff --git a/tests/pyathena/polars/test_result_set.py b/tests/pyathena/polars/test_result_set.py index 231f8daa8..2112af036 100644 --- a/tests/pyathena/polars/test_result_set.py +++ b/tests/pyathena/polars/test_result_set.py @@ -5,13 +5,19 @@ # # SPDX-License-Identifier: MIT +from datetime import datetime from unittest.mock import PropertyMock, patch import polars as pl import pytest from pyathena.error import OperationalError -from pyathena.polars.result_set import AthenaPolarsResultSet, PolarsDataFrameIterator +from pyathena.polars.converter import DefaultPolarsTypeConverter +from pyathena.polars.result_set import ( + AthenaPolarsResultSet, + PolarsDataFrameIterator, + _to_datetimes, +) _ROWS_BEFORE_FAILURE = 300_000 @@ -25,9 +31,48 @@ def _chunked_result_set() -> AthenaPolarsResultSet: result_set = AthenaPolarsResultSet.__new__(AthenaPolarsResultSet) # bypass __init__ result_set._chunksize = 10_000 result_set._kwargs = {} + result_set._metadata = None return result_set +@pytest.mark.parametrize( + ("dtype", "microseconds"), + [ + (pl.Datetime, [0, 123000, 123456, 123456, 123456]), + (pl.Datetime("ms"), [0, 123000, 123000, 123000, 123000]), + (pl.Datetime("us"), [0, 123000, 123456, 123456, 123456]), + ], +) +def test_to_datetimes(dtype, microseconds): + """Timestamp text with up to 12 fractional digits is truncated to the time unit. + + NULL can be null or an empty string, depending on the read options. A column that + the DataFrame does not have, such as one not selected, is skipped. + """ + df = pl.DataFrame( + { + "t": [ + "2020-01-02 03:04:05", + "2020-01-02 03:04:05.123", + "2020-01-02 03:04:05.123456", + "2020-01-02 03:04:05.123456789", + "2020-01-02 03:04:05.123456789012", + "0001-01-01 00:00:00.000", + "", + None, + ] + } + ) + result = _to_datetimes(df, {"t": dtype, "missing": dtype}) + assert result.schema["t"] == dtype + assert result["t"].to_list() == [ + *(datetime(2020, 1, 2, 3, 4, 5, us) for us in microseconds), + datetime(1, 1, 1), + None, + None, + ] + + class TestAthenaPolarsResultSet: def test_iter_csv_chunks_raises_when_read_fails_partway(self, tmp_path): """A CSV read that fails partway through the data raises instead of ending early.""" @@ -206,6 +251,51 @@ def test_storage_options_replace_defaults(self, reader, function): list(result) assert read.call_args.kwargs["storage_options"] == {"anon": True} + @pytest.mark.parametrize( + ("kwargs", "expected"), + [ + ({}, {"t": [datetime(2020, 1, 2, 3, 4, 5, 123456), None]}), + ({"columns": ["v"]}, {}), + ({"new_columns": ["t2", "v2"]}, {"t2": [datetime(2020, 1, 2, 3, 4, 5, 123456), None]}), + ({"with_column_names": lambda names: [n.upper() for n in names]}, OperationalError), + ], + ) + def test_read_csv_truncates_timestamps(self, kwargs, expected): + """Timestamps that fail to parse are read again as text and truncated. + + With ``with_column_names`` given to execute(), Polars renames the columns as it + reads them, so they are not, and the read fails as before. No AWS calls; the + GetQueryResults rows are mocked. + """ + result_set = AthenaPolarsResultSet.__new__(AthenaPolarsResultSet) # bypass __init__ + result_set._query_execution = None + result_set._converter = DefaultPolarsTypeConverter() + result_set._kwargs = kwargs + result_set._metadata = tuple( + {"Name": n, "Type": t, "Precision": 3, "Scale": 0, "Nullable": "UNKNOWN"} + for n, t in (("t", "timestamp"), ("v", "varchar")) + ) + data = b'"t","v"\n"2020-01-02 03:04:05.123456789012","x"\n,"y"\n' + with ( + patch.object(AthenaPolarsResultSet, "_fetch_all_rows_as_csv", return_value=data), + patch.object( + AthenaPolarsResultSet, + "_csv_storage_options", + new_callable=PropertyMock, + return_value={}, + ), + ): + if expected is OperationalError: + with pytest.raises(OperationalError): + result_set._read_csv() + return + df = result_set._read_csv() + for name, values in expected.items(): + assert df.schema[name] == pl.Datetime("us") + assert df[name].to_list() == values + if not expected: + assert df.columns == ["v"] + class TestPolarsDataFrameIterator: @pytest.mark.parametrize( diff --git a/tests/pyathena/test_converter.py b/tests/pyathena/test_converter.py index 9293047ff..fb32ba764 100644 --- a/tests/pyathena/test_converter.py +++ b/tests/pyathena/test_converter.py @@ -5,7 +5,6 @@ from pyathena.converter import ( DefaultTypeConverter, - _text_value_converter, _to_array, _to_datetime, _to_datetime_with_tz, @@ -675,21 +674,6 @@ def test_typed_time_with_tz_elements(): ) == {"a": time(12, 34, 56, tzinfo=jst)} -def test_text_value_converter(): - """The fallback converter keeps text values and nested time zones as text.""" - converter = _text_value_converter() - assert converter.convert("json", '{"a": 1}') == '{"a": 1}' - assert converter.convert("time with time zone", "12:34:56+09:00") == "12:34:56+09:00" - assert converter.convert( - "array", "[12:34:56.789+09:00]", type_hint="array(time with time zone)" - ) == ["12:34:56.789+09:00"] - assert converter.convert("array", '[{"a": 1}]', type_hint="array(json)") == [{"a": 1}] - # Other converters keep the default mappings. - assert DefaultTypeConverter().convert( - "array", "[12:34:56.789+09:00]", type_hint="array(time with time zone)" - ) == [time(12, 34, 56, 789000, tzinfo=timezone(timedelta(hours=9)))] - - @pytest.mark.parametrize( ("input_value", "expected"), [ diff --git a/tests/pyathena/test_result_set.py b/tests/pyathena/test_result_set.py index 4eb76777b..27d7c3b1a 100644 --- a/tests/pyathena/test_result_set.py +++ b/tests/pyathena/test_result_set.py @@ -23,6 +23,10 @@ def _page(values, next_token=None): return response +def _row(*values): + return {"Data": [{} if v is None else {"VarCharValue": v} for v in values]} + + class TestAthenaResultSet: def test_fetch_all_rows_skips_column_labels_only_on_first_page(self): """A later page can start with a data row equal to the column labels. @@ -44,6 +48,26 @@ def test_fetch_all_rows_skips_column_labels_only_on_first_page(self): with patch.object(result_set, "_get_query_results", side_effect=pages): assert result_set._fetch_all_rows() == [("1",), ("a",), ("2",)] + def test_fetch_all_rows_as_csv(self): + """The GetQueryResults rows are written as Athena writes a CSV result file. + + Only the first page starts with the column labels, so a later row equal to + them is data. No AWS calls; GetQueryResults is mocked. + """ + result_set = AthenaResultSet.__new__(AthenaResultSet) # bypass __init__ + result_set._query_execution = None + result_set._metadata = tuple( + {"Name": name, "Type": "varchar", "Precision": 0, "Scale": 0, "Nullable": "UNKNOWN"} + for name in ("v", 'w"') + ) + pages = [ + {"ResultSet": {"Rows": [_row("v", 'w"'), _row('a,"b"\nc', None)]}, "NextToken": "t"}, + {"ResultSet": {"Rows": [_row("v", 'w"'), _row("", "1")]}}, + ] + with patch.object(AthenaResultSet, "_get_query_results", side_effect=pages): + data = result_set._fetch_all_rows_as_csv() + assert data == b'"v","w"""\n"a,""b""\nc",\n"v","w"""\n"","1"\n' + class TestWithResultSet: def test_is_mixin(self): diff --git a/tests/pyathena/util.py b/tests/pyathena/util.py index cdd777999..cb1949640 100644 --- a/tests/pyathena/util.py +++ b/tests/pyathena/util.py @@ -144,6 +144,21 @@ def unreachable_glue(connection): None, ) +# Values that the cursors reading CSV result files must read the same way from the +# GetQueryResults rows of managed query result storage. +RESULT_FILE_VALUES_QUERY = """ +SELECT * FROM (VALUES + (1, 'a,"b"' || chr(10) || 'c', '', DATE '2020-01-02', TIMESTAMP '2020-01-02 03:04:05.123', + CAST('2020-01-02 03:04:05.123456' AS TIMESTAMP(6)), CAST('12.30' AS DECIMAL(10, 2)), X'0102', + ARRAY[1, 2], MAP(ARRAY['k'], ARRAY[1]), CAST(ROW(1, 'x') AS ROW(a INTEGER, b VARCHAR)), + json_parse('{"a": [1, 2]}'), BIGINT '9007199254740993', + CAST('2020-01-02 03:04:05.123456789012' AS TIMESTAMP(12))), + (2, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL, NULL) +) AS t(col_int, col_varchar, col_empty, col_date, col_timestamp, col_timestamp_6, col_decimal, + col_varbinary, col_array, col_map, col_row, col_json, col_bigint, col_timestamp_12) +ORDER BY col_int +""" + # A Spark job whose executor tasks sleep, so that StopCalculationExecution can cancel it. CANCELABLE_SPARK_JOB = """ From 1878a8c8b4349acadb61f0f390325537aeeb1466 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 14:26:40 +0900 Subject: [PATCH 2/3] Cover Arrow timestamp columns with the same name read as text Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/arrow/test_result_set.py | 24 ++++++++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/tests/pyathena/arrow/test_result_set.py b/tests/pyathena/arrow/test_result_set.py index 62f571486..777136064 100644 --- a/tests/pyathena/arrow/test_result_set.py +++ b/tests/pyathena/arrow/test_result_set.py @@ -71,3 +71,27 @@ def test_fetch_after_close(self): assert result_set.fetchone() is None assert result_set.fetchmany() == [] assert result_set.fetchall() == [] + + def test_read_csv_timestamps_with_the_same_name(self): + """Timestamp columns are converted by position, also with names that repeat. + + No AWS calls; the GetQueryResults rows are mocked. + """ + result_set = AthenaArrowResultSet.__new__(AthenaArrowResultSet) # bypass __init__ + result_set._query_execution = None + result_set._converter = DefaultArrowTypeConverter() + result_set._block_size = AthenaArrowResultSet.DEFAULT_BLOCK_SIZE + result_set._metadata = tuple( + {"Name": n, "Type": t, "Precision": 3, "Scale": 0, "Nullable": "UNKNOWN"} + for n, t in (("x", "varchar"), ("x", "timestamp"), ("y", "timestamp")) + ) + data = b'"x","x","y"\n"a","2020-01-02 03:04:05.123456789","2020-01-02 03:04:05.1"\n"b",,\n' + with patch.object(AthenaArrowResultSet, "_fetch_all_rows_as_csv", return_value=data): + table = result_set._read_csv() + assert table.column_names == ["x", "x", "y"] + assert table.schema.types == [pa.string(), pa.timestamp("us"), pa.timestamp("us")] + rows = list(zip(*(column.to_pylist() for column in table.columns), strict=True)) + assert rows == [ + ("a", datetime(2020, 1, 2, 3, 4, 5, 123456), datetime(2020, 1, 2, 3, 4, 5, 100000)), + ("b", None, None), + ] From d376f514df982203ff30337d222281ea9de5161a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 15:08:52 +0900 Subject: [PATCH 3/3] Move the timestamp text conversions to the converter modules _to_timestamp() and _to_datetimes() convert values, so they belong with the other conversions in the Arrow and Polars converter modules rather than in the result sets. The text lengths for each time unit are shared in pyathena.converter. Co-Authored-By: Claude Opus 5.5 --- pyathena/arrow/converter.py | 28 +++++++++++++- pyathena/arrow/result_set.py | 29 +-------------- pyathena/converter.py | 5 +++ pyathena/polars/converter.py | 38 ++++++++++++++++++- pyathena/polars/result_set.py | 38 +------------------ tests/pyathena/arrow/test_converter.py | 47 +++++++++++++++++++++++- tests/pyathena/arrow/test_result_set.py | 43 +--------------------- tests/pyathena/polars/test_converter.py | 45 ++++++++++++++++++++++- tests/pyathena/polars/test_result_set.py | 44 +--------------------- 9 files changed, 164 insertions(+), 153 deletions(-) diff --git a/pyathena/arrow/converter.py b/pyathena/arrow/converter.py index 51efec2fa..13d63caf8 100644 --- a/pyathena/arrow/converter.py +++ b/pyathena/arrow/converter.py @@ -5,9 +5,10 @@ import logging from collections.abc import Callable from copy import deepcopy -from typing import Any +from typing import TYPE_CHECKING, Any from pyathena.converter import ( + _TIMESTAMP_TEXT_LENGTHS, Converter, _to_binary, _to_date, @@ -20,6 +21,9 @@ ) from pyathena.util import override +if TYPE_CHECKING: + from pyarrow import ChunkedArray, TimestampType + _logger = logging.getLogger(__name__) @@ -34,6 +38,28 @@ } +def _to_timestamp(column: ChunkedArray, type_: TimestampType) -> ChunkedArray: + """Convert timestamp text to a timestamp type, truncating finer fractions. + + Athena writes up to 12 fractional digits, which pyarrow does not parse into a + timestamp type whose unit holds fewer. + + Args: + column: The timestamp text, with NULL as null or as an empty string. + type_: The timestamp type. + + Returns: + The timestamps. + """ + import pyarrow as pa + import pyarrow.compute as pc + + length = _TIMESTAMP_TEXT_LENGTHS[type_.unit] + if (pc.max(pc.utf8_length(column)).as_py() or 0) > length: + column = pc.utf8_slice_codeunits(column, 0, length) + return pc.if_else(pc.equal(column, ""), pa.scalar(None, pa.string()), column).cast(type_) + + class DefaultArrowTypeConverter(Converter): """Optimized type converter for Apache Arrow Table results. diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index 97353d5a1..2b1b9cbb9 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -11,6 +11,7 @@ ) from pyathena import OperationalError +from pyathena.arrow.converter import _to_timestamp from pyathena.arrow.util import to_column_info from pyathena.converter import Converter, _to_default from pyathena.model import AthenaQueryExecution @@ -19,38 +20,12 @@ if TYPE_CHECKING: import polars as pl - from pyarrow import ChunkedArray, Table, TimestampType + from pyarrow import Table from pyathena.connection import Connection _logger = logging.getLogger(__name__) -# The length of timestamp text, such as "2020-01-02 03:04:05.123456", that holds the -# fraction a timestamp unit can represent. -_TIMESTAMP_TEXT_LENGTHS: dict[str, int] = {"s": 19, "ms": 23, "us": 26, "ns": 29} - - -def _to_timestamp(column: ChunkedArray, type_: TimestampType) -> ChunkedArray: - """Convert timestamp text to a timestamp type, truncating finer fractions. - - Athena writes up to 12 fractional digits, which pyarrow does not parse into a - timestamp type whose unit holds fewer. - - Args: - column: The timestamp text, with NULL as null or as an empty string. - type_: The timestamp type. - - Returns: - The timestamps. - """ - import pyarrow as pa - import pyarrow.compute as pc - - length = _TIMESTAMP_TEXT_LENGTHS[type_.unit] - if (pc.max(pc.utf8_length(column)).as_py() or 0) > length: - column = pc.utf8_slice_codeunits(column, 0, length) - return pc.if_else(pc.equal(column, ""), pa.scalar(None, pa.string()), column).cast(type_) - class AthenaArrowResultSet(AthenaResultSet): """Result set that provides Apache Arrow Table results with columnar optimization. diff --git a/pyathena/converter.py b/pyathena/converter.py index a12513c37..644a50c00 100644 --- a/pyathena/converter.py +++ b/pyathena/converter.py @@ -68,6 +68,11 @@ def _to_datetime(varchar_value: str | None) -> datetime | None: return _parse_datetime(varchar_value) +# The length of Athena timestamp text, such as "2020-01-02 03:04:05.123456", that holds +# the fraction a time unit can represent. Athena writes up to 12 fractional digits. +_TIMESTAMP_TEXT_LENGTHS: dict[str, int] = {"s": 19, "ms": 23, "us": 26, "ns": 29} + + _UTC_OFFSET_PATTERN: re.Pattern[str] = re.compile(r"([+-])(\d{2}):(\d{2})") diff --git a/pyathena/polars/converter.py b/pyathena/polars/converter.py index 4d46c702f..a40e9c7d5 100644 --- a/pyathena/polars/converter.py +++ b/pyathena/polars/converter.py @@ -12,9 +12,10 @@ import logging from collections.abc import Callable from copy import deepcopy -from typing import Any +from typing import TYPE_CHECKING, Any from pyathena.converter import ( + _TIMESTAMP_TEXT_LENGTHS, Converter, _to_binary, _to_date, @@ -26,6 +27,9 @@ ) from pyathena.util import override +if TYPE_CHECKING: + import polars as pl + _logger = logging.getLogger(__name__) @@ -39,6 +43,38 @@ } +def _to_datetimes(df: pl.DataFrame, dtypes: dict[str, Any]) -> pl.DataFrame: + """Convert timestamp text columns to Datetime dtypes, truncating finer fractions. + + Athena writes up to 12 fractional digits, which Polars does not parse into a + Datetime whose time unit holds fewer. + + Args: + df: The DataFrame with the timestamp text columns, with NULL as null or as + an empty string. + dtypes: The Datetime dtypes keyed by column name. + + Returns: + The DataFrame with the columns converted. + """ + import polars as pl + + exprs = [] + for name, dtype in dtypes.items(): + if name not in df.columns: + # Not read, such as a column that the read options leave out. + continue + time_unit = (dtype() if isinstance(dtype, type) else dtype).time_unit + text = pl.col(name) + exprs.append( + pl.when(text != "") + .then(text.str.slice(0, _TIMESTAMP_TEXT_LENGTHS[time_unit])) + .str.to_datetime("%Y-%m-%d %H:%M:%S%.f", time_unit=time_unit) + .cast(dtype) + ) + return df.with_columns(exprs) if exprs else df + + class DefaultPolarsTypeConverter(Converter): """Optimized type converter for Polars DataFrame results. diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index 4ef0c750f..1ee5bc0bc 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -25,6 +25,7 @@ from pyathena.converter import Converter from pyathena.error import ProgrammingError from pyathena.model import AthenaQueryExecution +from pyathena.polars.converter import _to_datetimes from pyathena.polars.util import to_column_info from pyathena.result_set import AthenaResultSet from pyathena.util import RetryConfig, override @@ -38,48 +39,11 @@ _logger = logging.getLogger(__name__) -# The length of timestamp text, such as "2020-01-02 03:04:05.123456", that holds the -# fraction a Datetime time unit can represent. -_TIMESTAMP_TEXT_LENGTHS: dict[str, int] = {"ms": 23, "us": 26, "ns": 29} - - def _identity(x: Any) -> Any: """Identity function for use as default converter.""" return x -def _to_datetimes(df: pl.DataFrame, dtypes: dict[str, Any]) -> pl.DataFrame: - """Convert timestamp text columns to Datetime dtypes, truncating finer fractions. - - Athena writes up to 12 fractional digits, which Polars does not parse into a - Datetime whose time unit holds fewer. - - Args: - df: The DataFrame with the timestamp text columns, with NULL as null or as - an empty string. - dtypes: The Datetime dtypes keyed by column name. - - Returns: - The DataFrame with the columns converted. - """ - import polars as pl - - exprs = [] - for name, dtype in dtypes.items(): - if name not in df.columns: - # Not selected, as with ``columns`` given to ``execute()``. - continue - time_unit = (dtype() if isinstance(dtype, type) else dtype).time_unit - text = pl.col(name) - exprs.append( - pl.when(text != "") - .then(text.str.slice(0, _TIMESTAMP_TEXT_LENGTHS[time_unit])) - .str.to_datetime("%Y-%m-%d %H:%M:%S%.f", time_unit=time_unit) - .cast(dtype) - ) - return df.with_columns(exprs) if exprs else df - - class PolarsDataFrameIterator(abc.Iterator): # type: ignore[type-arg] """Iterator for chunked DataFrame results from Athena queries. diff --git a/tests/pyathena/arrow/test_converter.py b/tests/pyathena/arrow/test_converter.py index ae9c8cd90..5875e59ce 100644 --- a/tests/pyathena/arrow/test_converter.py +++ b/tests/pyathena/arrow/test_converter.py @@ -5,7 +5,12 @@ # # SPDX-License-Identifier: MIT -from pyathena.arrow.converter import DefaultArrowUnloadTypeConverter +from datetime import datetime + +import pyarrow as pa +import pytest + +from pyathena.arrow.converter import DefaultArrowUnloadTypeConverter, _to_timestamp class TestDefaultArrowUnloadTypeConverter: @@ -13,3 +18,43 @@ def test_convert_delegates_to_default(self): """convert() dispatches through the default converter instead of returning None.""" converter = DefaultArrowUnloadTypeConverter() assert converter.convert("varchar", "hello") == "hello" + + +@pytest.mark.parametrize( + ("unit", "microseconds"), + [ + ("s", [0, 0, 0, 0, 0]), + ("ms", [0, 123000, 123000, 123000, 123000]), + ("us", [0, 123000, 123456, 123456, 123456]), + ], +) +def test_to_timestamp(unit, microseconds): + """Timestamp text with up to 12 fractional digits is truncated to the unit. + + NULL can be null or an empty string, depending on the CSV read options. + """ + column = pa.chunked_array( + [ + pa.array( + [ + "2020-01-02 03:04:05", + "2020-01-02 03:04:05.123", + "2020-01-02 03:04:05.123456", + "2020-01-02 03:04:05.123456789", + "2020-01-02 03:04:05.123456789012", + "0001-01-01 00:00:00.000", + "", + None, + ], + pa.string(), + ) + ] + ) + values = _to_timestamp(column, pa.timestamp(unit)) + assert values.type == pa.timestamp(unit) + assert values.to_pylist() == [ + *(datetime(2020, 1, 2, 3, 4, 5, us) for us in microseconds), + datetime(1, 1, 1), + None, + None, + ] diff --git a/tests/pyathena/arrow/test_result_set.py b/tests/pyathena/arrow/test_result_set.py index 777136064..90cf33c9a 100644 --- a/tests/pyathena/arrow/test_result_set.py +++ b/tests/pyathena/arrow/test_result_set.py @@ -8,54 +8,13 @@ from unittest.mock import MagicMock, patch import pyarrow as pa -import pytest from pyathena.arrow.converter import DefaultArrowTypeConverter -from pyathena.arrow.result_set import AthenaArrowResultSet, _to_timestamp +from pyathena.arrow.result_set import AthenaArrowResultSet from pyathena.model import AthenaQueryExecution from pyathena.util import RetryConfig -@pytest.mark.parametrize( - ("unit", "microseconds"), - [ - ("s", [0, 0, 0, 0, 0]), - ("ms", [0, 123000, 123000, 123000, 123000]), - ("us", [0, 123000, 123456, 123456, 123456]), - ], -) -def test_to_timestamp(unit, microseconds): - """Timestamp text with up to 12 fractional digits is truncated to the unit. - - NULL can be null or an empty string, depending on the CSV read options. - """ - column = pa.chunked_array( - [ - pa.array( - [ - "2020-01-02 03:04:05", - "2020-01-02 03:04:05.123", - "2020-01-02 03:04:05.123456", - "2020-01-02 03:04:05.123456789", - "2020-01-02 03:04:05.123456789012", - "0001-01-01 00:00:00.000", - "", - None, - ], - pa.string(), - ) - ] - ) - values = _to_timestamp(column, pa.timestamp(unit)) - assert values.type == pa.timestamp(unit) - assert values.to_pylist() == [ - *(datetime(2020, 1, 2, 3, 4, 5, us) for us in microseconds), - datetime(1, 1, 1), - None, - None, - ] - - class TestAthenaArrowResultSet: def test_fetch_after_close(self): """No AWS calls; the query execution and the filesystem are mocked.""" diff --git a/tests/pyathena/polars/test_converter.py b/tests/pyathena/polars/test_converter.py index 07b1b7de0..f025e001d 100644 --- a/tests/pyathena/polars/test_converter.py +++ b/tests/pyathena/polars/test_converter.py @@ -5,7 +5,12 @@ # # SPDX-License-Identifier: MIT -from pyathena.polars.converter import DefaultPolarsUnloadTypeConverter +from datetime import datetime + +import polars as pl +import pytest + +from pyathena.polars.converter import DefaultPolarsUnloadTypeConverter, _to_datetimes class TestDefaultPolarsUnloadTypeConverter: @@ -13,3 +18,41 @@ def test_convert_delegates_to_default(self): """convert() dispatches through the default converter instead of returning None.""" converter = DefaultPolarsUnloadTypeConverter() assert converter.convert("varchar", "hello") == "hello" + + +@pytest.mark.parametrize( + ("dtype", "microseconds"), + [ + (pl.Datetime, [0, 123000, 123456, 123456, 123456]), + (pl.Datetime("ms"), [0, 123000, 123000, 123000, 123000]), + (pl.Datetime("us"), [0, 123000, 123456, 123456, 123456]), + ], +) +def test_to_datetimes(dtype, microseconds): + """Timestamp text with up to 12 fractional digits is truncated to the time unit. + + NULL can be null or an empty string, depending on the read options. A column that + the DataFrame does not have, such as one not selected, is skipped. + """ + df = pl.DataFrame( + { + "t": [ + "2020-01-02 03:04:05", + "2020-01-02 03:04:05.123", + "2020-01-02 03:04:05.123456", + "2020-01-02 03:04:05.123456789", + "2020-01-02 03:04:05.123456789012", + "0001-01-01 00:00:00.000", + "", + None, + ] + } + ) + result = _to_datetimes(df, {"t": dtype, "missing": dtype}) + assert result.schema["t"] == dtype + assert result["t"].to_list() == [ + *(datetime(2020, 1, 2, 3, 4, 5, us) for us in microseconds), + datetime(1, 1, 1), + None, + None, + ] diff --git a/tests/pyathena/polars/test_result_set.py b/tests/pyathena/polars/test_result_set.py index 2112af036..a39fbcb1a 100644 --- a/tests/pyathena/polars/test_result_set.py +++ b/tests/pyathena/polars/test_result_set.py @@ -13,11 +13,7 @@ from pyathena.error import OperationalError from pyathena.polars.converter import DefaultPolarsTypeConverter -from pyathena.polars.result_set import ( - AthenaPolarsResultSet, - PolarsDataFrameIterator, - _to_datetimes, -) +from pyathena.polars.result_set import AthenaPolarsResultSet, PolarsDataFrameIterator _ROWS_BEFORE_FAILURE = 300_000 @@ -35,44 +31,6 @@ def _chunked_result_set() -> AthenaPolarsResultSet: return result_set -@pytest.mark.parametrize( - ("dtype", "microseconds"), - [ - (pl.Datetime, [0, 123000, 123456, 123456, 123456]), - (pl.Datetime("ms"), [0, 123000, 123000, 123000, 123000]), - (pl.Datetime("us"), [0, 123000, 123456, 123456, 123456]), - ], -) -def test_to_datetimes(dtype, microseconds): - """Timestamp text with up to 12 fractional digits is truncated to the time unit. - - NULL can be null or an empty string, depending on the read options. A column that - the DataFrame does not have, such as one not selected, is skipped. - """ - df = pl.DataFrame( - { - "t": [ - "2020-01-02 03:04:05", - "2020-01-02 03:04:05.123", - "2020-01-02 03:04:05.123456", - "2020-01-02 03:04:05.123456789", - "2020-01-02 03:04:05.123456789012", - "0001-01-01 00:00:00.000", - "", - None, - ] - } - ) - result = _to_datetimes(df, {"t": dtype, "missing": dtype}) - assert result.schema["t"] == dtype - assert result["t"].to_list() == [ - *(datetime(2020, 1, 2, 3, 4, 5, us) for us in microseconds), - datetime(1, 1, 1), - None, - None, - ] - - class TestAthenaPolarsResultSet: def test_iter_csv_chunks_raises_when_read_fails_partway(self, tmp_path): """A CSV read that fails partway through the data raises instead of ending early."""