diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index cb12f808..8fb6c041 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -12,7 +12,7 @@ from pyathena import OperationalError from pyathena.arrow.util import to_column_info -from pyathena.converter import Converter, _text_value_converter, _to_default +from pyathena.converter import _TEXT_VALUE_TYPES, Converter, _text_value_converter, _to_default from pyathena.error import ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.result_set import AthenaResultSet @@ -254,23 +254,23 @@ def _fetch(self) -> None: except StopIteration: return else: - dict_rows = rows.to_pydict() - converters = ( - self.converters - if self._convert_rows - else self._text_value_converters(self.converters) - ) - if converters: - column_names = dict_rows.keys() + # Read the columns and their converters by position; to_pydict() and the + # 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 + ] + if any(convert is not _to_default for convert in converters): processed_rows = [ - tuple( - converters.get(k, _to_default)(v) - for k, v in zip(column_names, row, strict=False) - ) - for row in zip(*dict_rows.values(), strict=False) + tuple(convert(v) for convert, v in zip(converters, row, strict=False)) + for row in zip(*columns, strict=False) ] else: - processed_rows = list(zip(*dict_rows.values(), strict=False)) + processed_rows = list(zip(*columns, strict=False)) self._rows.extend(processed_rows) @override @@ -403,8 +403,8 @@ def _as_arrow_from_api(self, converter: Converter | None = None) -> Table: if not rows: return pa.Table.from_pydict({}) description = self.description if self.description else [] - columns = [d[0] for d in description] - return pa.table(self._rows_to_columnar(rows, columns)) + 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. @@ -447,4 +447,4 @@ def close(self) -> None: super().close() self._table = pa.Table.from_pydict({}) - self._batches = [] + self._batches = iter([]) diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index d700fa72..2112773b 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -411,9 +411,10 @@ def __init__( # The converters that pandas.read_csv() applies, keyed by column name. self._csv_converters: dict[Any, Callable[[str | None], Any]] = {} - # Cache time column names for efficient _trunc_date processing + # Cache time column names for efficient _trunc_date processing. _read_csv() + # replaces them with the labels of the columns that it reads. description = self.description if self.description else [] - self._time_columns: list[str] = [d[0] for d in description if d[1] == "time"] + self._time_columns: list[Any] = [d[0] for d in description if d[1] == "time"] import pandas as pd @@ -475,6 +476,12 @@ def _get_csv_engine( effective_chunksize is None and self._quoting == 1 and not self.converters + # The pyarrow engine does not rename columns with the same name, and + # the column labels cannot be resolved for it (see + # _get_csv_column_labels()). + and not self._needs_csv_column_name_resolution( + [d[0] for d in self.description or []] + ) and (file_size_bytes is None or file_size_bytes >= self.PYARROW_MIN_FILE_SIZE_BYTES) ) if is_compatible: @@ -588,6 +595,23 @@ def parse_dates(self) -> list[Any | None]: description = self.description if self.description else [] return [d[0] for d in description if d[1] in self._PARSE_DATES] + def _get_column_names(self) -> list[Any]: + """Get the names of the result columns in the DataFrame. + + Columns with the same name are renamed as pandas renames them when it reads + the header of a CSV file, such as ``x`` and ``x.1``. + + Returns: + List of column names. + """ + import pandas as pd + + description = self.description if self.description else [] + names = [d[0] for d in description] + if len(set(names)) == len(names): + return names + return self._resolve_csv_column_names(names, {}, pd.read_csv)[0] + def _finish_csv_frame(self, df: DataFrame) -> DataFrame: """Finish a DataFrame read from the CSV result file. @@ -641,8 +665,8 @@ def fetchone( return None else: self._rownumber = row[0] + 1 - description = self.description if self.description else [] - return tuple([row[1][d[0]] for d in description]) + # By position, so that columns with the same name keep their own values. + return tuple(row[1].values()) def _read_csv(self) -> TextFileReader | DataFrame: import pandas as pd @@ -676,11 +700,14 @@ def _read_csv(self) -> TextFileReader | DataFrame: csv_engine = self._get_csv_engine(length, effective_chunksize) read_csv_kwargs = self._get_csv_read_options(csv_engine, effective_chunksize) + labels = self._get_csv_column_labels(csv_engine, read_csv_kwargs) + if labels is not None: + self._key_csv_columns_by_labels(read_csv_kwargs, labels) try: with ExitStack() as stack: source: str | IOBase = self.output_location - binary_columns = self._configure_binary_csv_read(read_csv_kwargs, pd.read_csv) + binary_columns = self._configure_binary_csv_read(read_csv_kwargs, labels) self._csv_converters = read_csv_kwargs.get("converters") or {} if binary_columns: # Given storage_options, even None, open the file through fsspec @@ -721,7 +748,7 @@ def _get_csv_read_options(self, csv_engine: str, chunksize: int | None) -> dict[ if self.output_location and self.output_location.endswith(".txt"): sep = "\t" header = None - names = [d[0] for d in self.description or []] + names = self._get_column_names() else: sep = "," header = 0 @@ -796,12 +823,17 @@ def _resolve_csv_column_names( ) 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.""" + def _is_standard_csv_parsing(self, read_csv_kwargs: dict[str, Any]) -> bool: + """Whether the options read a CSV result file as Athena writes it. + + Args: + read_csv_kwargs: The options for ``pandas.read_csv()``. + + Returns: + True if pandas reads the header and quoted fields of the file as written. + """ return not ( - "varbinary" not in self._converter.mappings - or "converters" in self._kwargs - or not self.output_location + 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 @@ -810,6 +842,14 @@ def _can_preserve_binary_csv_nulls(self, read_csv_kwargs: dict[str, Any]) -> boo or read_csv_kwargs.get("quotechar", '"') != '"' ) + 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 ( + "varbinary" in self._converter.mappings + and "converters" not in self._kwargs + and self._is_standard_csv_parsing(read_csv_kwargs) + ) + def _needs_csv_column_name_resolution(self, column_names: list[Any]) -> bool: """Whether pandas must resolve column names instead of using Athena metadata.""" return ( @@ -829,38 +869,98 @@ def _needs_csv_column_name_resolution(self, column_names: list[Any]) -> bool: ) ) + def _get_csv_column_labels( + self, csv_engine: str, read_csv_kwargs: dict[str, Any] + ) -> list[Any] | None: + """Get the labels that pandas gives the result columns when it reads the CSV file. + + Args: + csv_engine: The CSV engine that reads the file. + read_csv_kwargs: The options for ``pandas.read_csv()``. + + Returns: + The label of each result column in the description order, with None for a + column that ``usecols`` leaves out. None if the options do not read the file + as Athena writes it, or if the labels do not match the result columns one + to one, such as with fewer ``names``. + """ + import pandas as pd + + # The pyarrow engine runs only when no labels need resolving (see + # _get_csv_engine()), and does not support reading only the header. + if csv_engine == "pyarrow" or not self._is_standard_csv_parsing(read_csv_kwargs): + return None + column_names = [d[0] for d in self.description or []] + if not self._needs_csv_column_name_resolution(column_names): + return column_names + labels, selected_labels = self._resolve_csv_column_names( + column_names, read_csv_kwargs, pd.read_csv + ) + if len(labels) != len(column_names): + return None + return [label if label in selected_labels else None for label in labels] + + def _key_csv_columns_by_labels( + self, read_csv_kwargs: dict[str, Any], labels: list[Any] + ) -> None: + """Key the column options of ``pandas.read_csv()`` by the labels of the columns. + + Columns with the same name keep their own converters and date parsing, and + options that rename or select columns get the types of the columns they + read. The ``dtype``, ``converters``, and ``parse_dates`` given to + ``execute()`` are kept as they are. + + Args: + read_csv_kwargs: The options for ``pandas.read_csv()``, updated in place. + labels: The labels from ``_get_csv_column_labels()``. + """ + description = self.description or [] + columns = [ + (label, d) for label, d in zip(labels, description, strict=True) if label is not None + ] + if "dtype" not in self._kwargs: + read_csv_kwargs["dtype"] = { + label: dtype + for label, d in columns + if (dtype := self._converter.get_dtype(d[1], d[4], d[5])) is not None + } + if "converters" not in self._kwargs: + read_csv_kwargs["converters"] = { + label: self._get_csv_converter(d[1]) + for label, d in columns + if d[1] in self._converter.mappings + } + if "parse_dates" not in self._kwargs: + read_csv_kwargs["parse_dates"] = [ + label for label, d in columns if d[1] in self._PARSE_DATES + ] + self._time_columns = [label for label, d in columns if d[1] == "time"] + def _configure_binary_csv_read( - self, read_csv_kwargs: dict[str, Any], read_csv: Callable[..., DataFrame] + self, read_csv_kwargs: dict[str, Any], labels: list[Any] | None ) -> 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() + """Wrap binary converters and return column positions needing NULL preservation. - description = self.description or [] - binary_columns = {i for i, d in enumerate(description) if d[1] == "varbinary"} - if not binary_columns: - return set() + Args: + read_csv_kwargs: The options for ``pandas.read_csv()``, whose converters + are keyed by the column labels. + labels: The labels from ``_get_csv_column_labels()``. - 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._get_csv_converter(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} + Returns: + The positions of the binary columns whose NULL fields the stream preserves. + """ + if labels is None or 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" and labels[i] is not None + } if binary_columns: + converters = read_csv_kwargs["converters"] for index in binary_columns: - name = column_names[index] - converters[name] = partial(_convert_binary_csv, converters[name]) - read_csv_kwargs["converters"] = converters + label = labels[index] + converters[label] = partial(_convert_binary_csv, converters[label]) return binary_columns def _open_binary_csv_stream( @@ -957,24 +1057,20 @@ def _as_pandas_from_api(self, converter: Converter | None = None) -> DataFrame: if not rows: return pd.DataFrame() description = self.description if self.description else [] - columns = [d[0] for d in description] - columnar = self._rows_to_columnar(rows, columns) + # 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. - dtypes: dict[str, Any] = {} - for d in description: + 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: - if (dtype := self._converter.get_dtype(d[1], d[4], d[5])) is not None: - dtypes[d[0]] = dtype - elif d[1] == "json" and None in columnar[d[0]]: - dtypes[d[0]] = object - return pd.DataFrame( - { - name: values if name not in dtypes else pd.array(values, dtype=dtypes[name]) - for name, values in columnar.items() - } - ) + 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 7f9bea74..85ea4742 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -9,9 +9,11 @@ from __future__ import annotations +import csv import logging from collections import abc from collections.abc import Callable, Iterator +from io import BytesIO, StringIO from multiprocessing import cpu_count from typing import ( TYPE_CHECKING, @@ -272,27 +274,30 @@ def __init__( if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: if self._chunksize is None: self._df = self._as_polars() - self._df_converters = self.converters + self._df_converters = self._get_converters(self._get_frame_column_names()) 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. - self._df_converters = self._text_value_converters(self.converters) + column_names = self._get_frame_column_names() + self._df_converters = self._text_value_converters( + self._get_converters(column_names), column_names + ) else: self._df = pl.DataFrame() if self._df is not None: # A clone keeps assignments to the DataFrame from as_polars() # out of the rows that the fetch methods return. self._df_iter = PolarsDataFrameIterator( - self._df.clone(), self._df_converters, self._get_column_names() + self._df.clone(), self._df_converters, self._get_frame_column_names() ) # Cache column names for efficient access in fetchone() # Must be after _as_polars() and _create_dataframe_iterator(), which update # _metadata for unload - self._column_names_cache: list[str] = self._get_column_names() + self._column_names_cache: list[str] = self._get_frame_column_names() self._iterrows = self._df_iter.iterrows() def _storage_options(self, default: Callable[[], dict[str, Any]]) -> Any: @@ -377,11 +382,7 @@ def _parquet_storage_options(self) -> dict[str, Any]: def dtypes(self) -> dict[str, Any]: """Get Polars-compatible data types for result columns.""" description = self.description if self.description else [] - return { - d[0]: dtype - for d in description - if (dtype := self._converter.get_dtype(d[1], d[4], d[5])) is not None - } + return self._get_dtypes([d[0] for d in description]) @property def converters(self) -> dict[str, Callable[[str | None], Any | None]]: @@ -391,16 +392,79 @@ def converters(self) -> dict[str, Callable[[str | None], Any | None]]: Dictionary mapping column names to their converter functions. """ description = self.description if self.description else [] - return {d[0]: self._converter.get(d[1]) for d in description} + return self._get_converters([d[0] for d in description]) + + @property + def _csv_dtypes(self) -> dict[str, Any]: + """The Polars data types of the result columns, keyed by the header of a CSV file.""" + return self._get_dtypes(self._get_column_names()) + + def _get_dtypes(self, column_names: list[str]) -> dict[str, Any]: + """Get the Polars data types of the result columns. + + Args: + column_names: The names of the result columns, in column order. + + Returns: + The data types keyed by the given names. + """ + description = self.description if self.description else [] + return { + name: dtype + for name, d in zip(column_names, description, strict=True) + if (dtype := self._converter.get_dtype(d[1], d[4], d[5])) is not None + } + + def _get_converters( + self, column_names: list[str] + ) -> dict[str, Callable[[str | None], Any | None]]: + """Get the conversion functions of the result columns. + + Args: + column_names: The names of the result columns, in column order. + + Returns: + The conversion functions keyed by the given names. + """ + description = self.description if self.description else [] + return { + name: self._converter.get(d[1]) + for name, d in zip(column_names, description, strict=True) + } def _get_column_names(self) -> list[str]: - """Get column names from description. + """Get the names of the result columns in a DataFrame. + + Columns with the same name are renamed as Polars renames them when it reads + the header of a CSV file, such as ``x`` and ``x_duplicated_0``. Returns: List of column names. """ + import polars as pl + description = self.description if self.description else [] - return [d[0] for d in description] + names = [d[0] for d in description] + if len(set(names)) == len(names): + return names + header = StringIO() + csv.writer(header, quoting=csv.QUOTE_ALL).writerow(names) + return pl.read_csv(BytesIO(header.getvalue().encode()), n_rows=0).columns + + def _get_frame_column_names(self) -> list[str]: + """Get the names of the result columns in the DataFrames that the result set reads. + + The ``new_columns`` given to ``execute()`` rename the columns of a CSV result + file by position. + + Returns: + List of column names. + """ + 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: + return names + return [*new_columns[: len(names)], *names[len(new_columns) :]] def _create_dataframe_iterator(self) -> PolarsDataFrameIterator: """Create a DataFrame iterator that reads the result file in chunks. @@ -420,7 +484,8 @@ def _create_dataframe_iterator(self) -> PolarsDataFrameIterator: else: self._metadata = () reader = iter(()) - return PolarsDataFrameIterator(reader, self.converters, self._get_column_names()) + column_names = self._get_frame_column_names() + return PolarsDataFrameIterator(reader, self._get_converters(column_names), column_names) @override def fetchone( @@ -494,19 +559,20 @@ def _read_csv(self) -> pl.DataFrame: raise ProgrammingError("output_location is not available.") separator, has_header, new_columns = self._get_csv_params() + read_kwargs = self._read_kwargs( + lambda: self._csv_storage_options, + separator=separator, + has_header=has_header, + schema_overrides=self._csv_dtypes, + ) + 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) try: - df = pl.read_csv( - self.output_location, - **self._read_kwargs( - lambda: self._csv_storage_options, - separator=separator, - has_header=has_header, - schema_overrides=self.dtypes, - ), - ) + df = pl.read_csv(self.output_location, **read_kwargs) if new_columns: - df.columns = 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}.") @@ -591,9 +657,8 @@ def _as_polars_from_api(self, converter: Converter | None = None) -> pl.DataFram rows = self._fetch_all_rows(converter or _text_value_converter()) if not rows: return pl.DataFrame() - description = self.description if self.description else [] - columns = [d[0] for d in description] - return pl.DataFrame(self._rows_to_columnar(rows, columns)) + 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. @@ -654,16 +719,22 @@ def _get_csv_params(self) -> tuple[str, bool, list[str] | None]: """Get CSV parsing parameters based on file type. Returns: - Tuple of (separator, has_header, new_columns). + Tuple of (separator, has_header, new_columns). ``new_columns`` are the + names of the first columns, which the readers set after reading as Polars + sets the ``new_columns`` given to ``execute()``. With ``schema_overrides`` + given to ``execute()``, which replace the result set's types, Polars + renames the columns of a CSV file itself. """ if self.output_location and self.output_location.endswith(".txt"): separator = "\t" has_header = False - new_columns: list[str] | None = self._get_column_names() + new_columns: list[str] | None = self._get_frame_column_names() else: separator = "," has_header = True - new_columns = None + new_columns = ( + None if "schema_overrides" in self._kwargs else self._kwargs.get("new_columns") + ) return separator, has_header, new_columns def _iter_csv_chunks(self) -> Iterator[pl.DataFrame]: @@ -685,22 +756,23 @@ def _iter_csv_chunks(self) -> Iterator[pl.DataFrame]: raise ProgrammingError("output_location is not available.") separator, has_header, new_columns = self._get_csv_params() + # scan_csv uses Rust's native object_store (like scan_parquet), + # not fsspec, so we use the same storage options as Parquet + read_kwargs = self._read_kwargs( + lambda: self._parquet_storage_options, + separator=separator, + has_header=has_header, + schema_overrides=self._csv_dtypes, + ) + 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) try: - # scan_csv uses Rust's native object_store (like scan_parquet), - # not fsspec, so we use the same storage options as Parquet - lazy_df = pl.scan_csv( - self.output_location, - **self._read_kwargs( - lambda: self._parquet_storage_options, - separator=separator, - has_header=has_header, - schema_overrides=self.dtypes, - ), - ) + lazy_df = pl.scan_csv(self.output_location, **read_kwargs) for batch in lazy_df.collect_batches(chunk_size=self._chunksize): if new_columns: - batch.columns = new_columns + batch.columns = [*new_columns, *batch.columns[len(new_columns) :]] yield batch except Exception as e: _logger.exception(f"Failed to read {self.output_location}.") @@ -763,7 +835,9 @@ def iter_chunks(self) -> PolarsDataFrameIterator: ... process(df) # Single DataFrame with all data """ if self._df is not None: - return PolarsDataFrameIterator(self._df, self._df_converters, self._get_column_names()) + return PolarsDataFrameIterator( + self._df, self._df_converters, self._get_frame_column_names() + ) return self._df_iter @override diff --git a/pyathena/result_set.py b/pyathena/result_set.py index 9ecc9187..81f1167a 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -682,18 +682,28 @@ def _is_first_row_column_labels(self, rows: list[dict[str, Any]]) -> bool: return True def _text_value_converters( - self, converters: dict[str, Callable[[str | None], Any | None]] + 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. 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. Returns: The converters of the columns whose Athena type is in ``_TEXT_VALUE_TYPES``. """ description = self.description if self.description else [] - return {d[0]: converters[d[0]] for d in description if d[1] in _TEXT_VALUE_TYPES} + 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 + } def _fetch_all_rows( self, @@ -731,10 +741,12 @@ def _fetch_all_rows( 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) - offset = 1 if rows and self._is_first_row_column_labels(rows) else 0 + # 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 all_rows.extend( cast( list[tuple[Any | None, ...]], @@ -747,26 +759,6 @@ def _fetch_all_rows( return all_rows - @staticmethod - def _rows_to_columnar( - rows: list[tuple[Any | None, ...]], - columns: list[str], - ) -> dict[str, list[Any]]: - """Convert row-oriented data to columnar format. - - Args: - rows: List of row tuples from ``_fetch_all_rows()``. - columns: Column names in order. - - Returns: - Dictionary mapping column names to lists of values. - """ - columnar: dict[str, list[Any]] = {col: [] for col in columns} - for row in rows: - for col, val in zip(columns, row, strict=False): - columnar[col].append(val) - return columnar - def _get_content_length(self) -> int: if not self.output_location: raise ProgrammingError("OutputLocation is none or empty.") diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index f9a66712..be738f3a 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -1054,6 +1054,39 @@ def test_fetch_all_rows(self, arrow_cursor): "2024-02-29 23:59:58.123 +05:30" ] + @pytest.mark.parametrize( + "arrow_cursor", + [ + pytest.param({}, id="default"), + pytest.param( + {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}, + id="managed", + marks=pytest.mark.skipif( + not ENV.managed_work_group, + reason="AWS_ATHENA_MANAGED_WORKGROUP not set", + ), + ), + ], + indirect=["arrow_cursor"], + ) + def test_duplicate_column_names(self, arrow_cursor): + arrow_cursor.execute( + "SELECT 1 AS x, 2 AS x, 'a' AS y, json_parse('[1]') AS j, 'b' AS j, " + "CAST('12:34:56' AS TIME) AS t, CAST('01:02:03' AS TIME) AS t" + ) + assert arrow_cursor.fetchall() == [ + ( + 1, + 2, + "a", + [1], + "b", + datetime(2017, 1, 1, 12, 34, 56).time(), + datetime(2017, 1, 1, 1, 2, 3).time(), + ) + ] + assert arrow_cursor.as_arrow().column_names == ["x", "x", "y", "j", "j", "t", "t"] + @pytest.mark.parametrize( "execute_kwargs", [{}, {"connect_timeout": 3.0, "request_timeout": 4.0}] ) diff --git a/tests/pyathena/arrow/test_result_set.py b/tests/pyathena/arrow/test_result_set.py new file mode 100644 index 00000000..c79aa1e2 --- /dev/null +++ b/tests/pyathena/arrow/test_result_set.py @@ -0,0 +1,29 @@ +# Copyright 2026 The PyAthena authors +# +# Licensed under the MIT License. +# See LICENSE or https://opensource.org/licenses/MIT. +# +# SPDX-License-Identifier: MIT +from unittest.mock import MagicMock, patch + +from pyathena.arrow.converter import DefaultArrowTypeConverter +from pyathena.arrow.result_set import AthenaArrowResultSet +from pyathena.model import AthenaQueryExecution +from pyathena.util import RetryConfig + + +class TestAthenaArrowResultSet: + def test_fetch_after_close(self): + """No AWS calls; the query execution and the filesystem are mocked.""" + with patch.object(AthenaArrowResultSet, "_create_s3_file_system"): + result_set = AthenaArrowResultSet( + connection=MagicMock(), + converter=DefaultArrowTypeConverter(), + query_execution=MagicMock(state=AthenaQueryExecution.STATE_FAILED), + arraysize=1, + retry_config=RetryConfig(), + ) + result_set.close() + assert result_set.fetchone() is None + assert result_set.fetchmany() == [] + assert result_set.fetchall() == [] diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 9d0f1279..96fb6f23 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -967,6 +967,8 @@ def test_get_csv_engine_explicit_specification(self): result_set = AthenaPandasResultSet.__new__(AthenaPandasResultSet) result_set._chunksize = None # Default values result_set._quoting = 1 + result_set._metadata = None + result_set._kwargs = {} # Test C engine specification result_set._engine = "c" @@ -989,6 +991,34 @@ def test_get_csv_engine_explicit_specification(self): engine = result_set._get_csv_engine() assert engine == "pyarrow" + # Test PyArrow with column names that repeat, which it does not rename + with ( + patch.object(result_set, "_get_available_engine", return_value="pyarrow"), + patch.object( + type(result_set), "converters", new_callable=PropertyMock, return_value={} + ), + patch.object( + type(result_set), + "description", + new_callable=PropertyMock, + return_value=[("x", "integer"), ("x", "integer")], + ), + ): + engine = result_set._get_csv_engine() + assert engine == "c" + + # Test PyArrow with read options that rename the columns + result_set._kwargs = {"names": ["b", "a"]} + with ( + patch.object(result_set, "_get_available_engine", return_value="pyarrow"), + patch.object( + type(result_set), "converters", new_callable=PropertyMock, return_value={} + ), + ): + engine = result_set._get_csv_engine() + assert engine == "c" + result_set._kwargs = {} + # Test PyArrow with incompatible chunksize (via parameter) with ( patch.object( @@ -1637,6 +1667,63 @@ def test_fetch_all_rows(self, pandas_cursor): pandas_cursor.execute(CONVERTED_VALUES_QUERY) assert pandas_cursor.fetchall() == [CONVERTED_VALUES_ROW] + @pytest.mark.parametrize( + "pandas_cursor", + [ + pytest.param({}, id="default"), + pytest.param( + {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}, + id="managed", + marks=pytest.mark.skipif( + not ENV.managed_work_group, + reason="AWS_ATHENA_MANAGED_WORKGROUP not set", + ), + ), + ], + indirect=["pandas_cursor"], + ) + def test_duplicate_column_names(self, pandas_cursor): + pandas_cursor.execute( + "SELECT 1 AS x, 'a' AS x, 'b' AS y, json_parse('[1]') AS j, json_parse('[2]') AS j, " + "CAST('12:34:56' AS TIME) AS t, 2 AS t" + ) + assert pandas_cursor.fetchall() == [ + ( + 1, + "a", + "b", + [1], + [2], + datetime(2017, 1, 1, 12, 34, 56).time(), + 2, + ) + ] + assert pandas_cursor.as_pandas().columns.tolist() == [ + "x", + "x.1", + "y", + "j", + "j.1", + "t", + "t.1", + ] + + @pytest.mark.parametrize( + ("execute_kwargs", "expected_row", "expected_columns"), + [ + ({"names": ["a", "b", "x.1"]}, (1, 2, "c"), ["a", "b", "x.1"]), + ({"usecols": [1, 2]}, (2, "c"), ["x.1", "y"]), + ({"usecols": ["x.1", "y"]}, (2, "c"), ["x.1", "y"]), + ], + ) + def test_duplicate_column_names_read_options( + self, pandas_cursor, execute_kwargs, expected_row, expected_columns + ): + """The column types follow the columns that the read options rename or select.""" + pandas_cursor.execute("SELECT 1 AS x, 2 AS x, 'c' AS y", **execute_kwargs) + assert pandas_cursor.fetchall() == [expected_row] + assert pandas_cursor.as_pandas().columns.tolist() == expected_columns + @pytest.mark.parametrize( "execute_kwargs", [ diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index a943d268..d4c642eb 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -787,6 +787,68 @@ def test_fetch_all_rows(self, polars_cursor): "2024-02-29 23:59:58.123 +05:30" ] + @pytest.mark.parametrize( + "polars_cursor", + [ + pytest.param({}, id="default"), + pytest.param( + {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}, + id="managed", + marks=pytest.mark.skipif( + not ENV.managed_work_group, + reason="AWS_ATHENA_MANAGED_WORKGROUP not set", + ), + ), + ], + indirect=["polars_cursor"], + ) + def test_duplicate_column_names(self, polars_cursor): + polars_cursor.execute( + "SELECT 1 AS x, 'a' AS x, 'b' AS y, json_parse('[1]') AS j, json_parse('[2]') AS j, " + "CAST('12:34:56' AS TIME) AS t, 2 AS t" + ) + assert polars_cursor.fetchall() == [ + ( + 1, + "a", + "b", + [1], + [2], + datetime(2017, 1, 1, 12, 34, 56).time(), + 2, + ) + ] + assert polars_cursor.as_polars().columns == [ + "x", + "x_duplicated_0", + "y", + "j", + "j_duplicated_0", + "t", + "t_duplicated_0", + ] + + @pytest.mark.parametrize("chunksize", [None, 1]) + def test_new_columns_renaming_first_columns(self, polars_cursor, chunksize): + """The types stay with the columns when new_columns renames only the first ones.""" + polars_cursor.execute("SELECT '001' AS x, 2 AS y", new_columns=["z"], chunksize=chunksize) + assert polars_cursor.fetchall() == [("001", 2)] + + def test_new_columns_with_schema_overrides(self, polars_cursor): + """schema_overrides given with new_columns are keyed by the new names, as in Polars.""" + polars_cursor.execute( + "SELECT '001' AS x, 2 AS y", new_columns=["z"], schema_overrides={"z": pl.String} + ) + assert polars_cursor.fetchall() == [("001", 2)] + + def test_duplicate_column_names_new_columns(self, polars_cursor): + """The column types follow the columns that new_columns renames.""" + polars_cursor.execute( + "SELECT 1 AS x, 2 AS x, 'c' AS y", new_columns=["y", "x_duplicated_0", "x"] + ) + assert polars_cursor.fetchall() == [(1, 2, "c")] + assert polars_cursor.as_polars().columns == ["y", "x_duplicated_0", "x"] + @pytest.mark.parametrize( "execute_kwargs", [{}, {"block_size": 2048, "cache_type": "none", "max_workers": 3, "chunksize": 20}], diff --git a/tests/pyathena/polars/test_result_set.py b/tests/pyathena/polars/test_result_set.py index c16318ec..231f8daa 100644 --- a/tests/pyathena/polars/test_result_set.py +++ b/tests/pyathena/polars/test_result_set.py @@ -45,7 +45,7 @@ def test_iter_csv_chunks_raises_when_read_fails_partway(self, tmp_path): ), patch.object( AthenaPolarsResultSet, - "dtypes", + "_csv_dtypes", new_callable=PropertyMock, return_value={"a": pl.Int64}, ), @@ -101,7 +101,7 @@ def test_csv_read_kwargs_replace_defaults(self, tmp_path, reader): ), patch.object( AthenaPolarsResultSet, - "dtypes", + "_csv_dtypes", new_callable=PropertyMock, return_value={"1;x": pl.Int64}, ), @@ -123,6 +123,42 @@ def test_csv_read_kwargs_replace_defaults(self, tmp_path, reader): df = result if isinstance(result, pl.DataFrame) else pl.concat(list(result)) assert df.to_dict(as_series=False) == {"column_1": ["1", "2"], "column_2": ["x", "y"]} + @pytest.mark.parametrize("reader", ["_read_csv", "_iter_csv_chunks"]) + def test_txt_new_columns_with_schema_overrides(self, tmp_path, reader): + """schema_overrides given with new_columns reach Polars with them for a .txt file.""" + path = tmp_path / "result.txt" + path.write_text("001\tx\n") + result_set = _chunked_result_set() + result_set._kwargs = {"new_columns": ["z"], "schema_overrides": {"z": pl.Utf8}} + with ( + patch.object( + AthenaPolarsResultSet, + "output_location", + new_callable=PropertyMock, + return_value=str(path), + ), + patch.object(AthenaPolarsResultSet, "_get_column_names", return_value=["a", "b"]), + patch.object( + AthenaPolarsResultSet, "_csv_dtypes", new_callable=PropertyMock, return_value={} + ), + patch.object( + AthenaPolarsResultSet, + "_csv_storage_options", + new_callable=PropertyMock, + return_value={}, + ), + patch.object( + AthenaPolarsResultSet, + "_parquet_storage_options", + new_callable=PropertyMock, + return_value={}, + ), + patch.object(AthenaPolarsResultSet, "_is_csv_readable", return_value=True), + ): + result = getattr(result_set, reader)() + df = result if isinstance(result, pl.DataFrame) else pl.concat(list(result)) + assert df.to_dict(as_series=False) == {"z": ["001"], "b": ["x"]} + @pytest.mark.parametrize( ("reader", "function"), [ @@ -146,7 +182,7 @@ def test_storage_options_replace_defaults(self, reader, function): return_value="s3://bucket/result.csv", ), patch.object( - AthenaPolarsResultSet, "dtypes", new_callable=PropertyMock, return_value={} + AthenaPolarsResultSet, "_csv_dtypes", new_callable=PropertyMock, return_value={} ), patch.object( AthenaPolarsResultSet, diff --git a/tests/pyathena/test_result_set.py b/tests/pyathena/test_result_set.py index 26635f24..4eb76777 100644 --- a/tests/pyathena/test_result_set.py +++ b/tests/pyathena/test_result_set.py @@ -4,11 +4,45 @@ # See LICENSE or https://opensource.org/licenses/MIT. # # SPDX-License-Identifier: MIT +from unittest.mock import MagicMock, patch + import pytest from pyathena.aio.common import AioBaseCursor, WithAsyncFetch from pyathena.common import BaseCursor, CursorIterator -from pyathena.result_set import WithFetch, WithResultSet +from pyathena.converter import DefaultTypeConverter +from pyathena.model import AthenaQueryExecution +from pyathena.result_set import AthenaResultSet, WithFetch, WithResultSet +from pyathena.util import RetryConfig + + +def _page(values, next_token=None): + response = {"ResultSet": {"Rows": [{"Data": [{"VarCharValue": v}]} for v in values]}} + if next_token: + response["NextToken"] = next_token + return response + + +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. + + No AWS calls; the GetQueryResults pages are mocked. + """ + result_set = AthenaResultSet( + connection=MagicMock(), + converter=DefaultTypeConverter(), + query_execution=MagicMock(state=AthenaQueryExecution.STATE_SUCCEEDED), + arraysize=1, + retry_config=RetryConfig(), + _pre_fetch=False, + ) + result_set._process_metadata( + {"ResultSet": {"ResultSetMetadata": {"ColumnInfo": [{"Name": "a", "Type": "varchar"}]}}} + ) + pages = [_page(["a", "1"], "token"), _page(["a", "2"])] + with patch.object(result_set, "_get_query_results", side_effect=pages): + assert result_set._fetch_all_rows() == [("1",), ("a",), ("2",)] class TestWithResultSet: