From d7242877b8052662b88bee45b073fd3d2ad799e0 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 01:56:26 +0900 Subject: [PATCH 1/6] Keep columns with the same name and header-like rows on later fallback pages - Read the pandas and Polars result columns by the names their CSV readers give columns with the same name (x.1, x_duplicated_0), and build the GetQueryResults fallback DataFrames with the same names. - Read Arrow record batches by position instead of to_pydict(), and build the fallback table from positional columns. - Skip the column labels only on the first GetQueryResults page of the fallback. - Let fetching from a closed Arrow result set return no rows. Fixes #1032. Co-Authored-By: Claude Opus 5.5 --- pyathena/arrow/result_set.py | 22 ++++---- pyathena/pandas/result_set.py | 72 +++++++++++++++++-------- pyathena/polars/result_set.py | 34 ++++++++---- pyathena/result_set.py | 38 ++++++------- tests/pyathena/arrow/test_cursor.py | 33 ++++++++++++ tests/pyathena/arrow/test_result_set.py | 29 ++++++++++ tests/pyathena/pandas/test_cursor.py | 41 ++++++++++++++ tests/pyathena/polars/test_cursor.py | 41 ++++++++++++++ tests/pyathena/test_result_set.py | 36 ++++++++++++- 9 files changed, 280 insertions(+), 66 deletions(-) create mode 100644 tests/pyathena/arrow/test_result_set.py diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index cb12f808e..12d82e671 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -254,23 +254,23 @@ def _fetch(self) -> None: except StopIteration: return else: - dict_rows = rows.to_pydict() + # Read the columns by position; to_pydict() keeps one column per name. + columns = [column.to_pylist() for column in rows.columns] converters = ( self.converters if self._convert_rows else self._text_value_converters(self.converters) ) if converters: - column_names = dict_rows.keys() + column_converters = [ + converters.get(name, _to_default) for name in rows.schema.names + ] 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(column_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 d700fa720..3c583d0f5 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -413,7 +413,11 @@ def __init__( # Cache time column names for efficient _trunc_date processing 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[str] = [ + name + for name, d in zip(self._get_column_names(), description, strict=True) + if d[1] == "time" + ] import pandas as pd @@ -436,6 +440,9 @@ def __init__( # out of the rows that the fetch methods return. Mutable values in its # cells, such as lists from JSON columns, are still shared. self._df_iter = PandasDataFrameIterator(self._df.copy(deep=False), _no_trunc_date) + # Cache column names for fetchone(), after _as_pandas(), which replaces the + # metadata of unload queries. + self._column_names_cache = self._get_column_names() self._iterrows = self._df_iter.iterrows() def _get_parquet_engine(self) -> str: @@ -471,10 +478,13 @@ def _get_csv_engine( # checks pass; otherwise fall through to the C engine default. if self._engine == "pyarrow": effective_chunksize = chunksize if chunksize is not None else self._chunksize + column_names = [d[0] for d in self.description or []] is_compatible = ( effective_chunksize is None and self._quoting == 1 and not self.converters + # The pyarrow engine does not rename columns with the same name. + and len(set(column_names)) == len(column_names) and (file_size_bytes is None or file_size_bytes >= self.PYARROW_MIN_FILE_SIZE_BYTES) ) if is_compatible: @@ -567,8 +577,8 @@ def dtypes(self) -> dict[str, type[Any]]: """ description = self.description if self.description else [] return { - d[0]: dtype - for d in description + name: dtype + for name, d in zip(self._get_column_names(), description, strict=True) if (dtype := self._converter.get_dtype(d[1], d[4], d[5])) is not None } @@ -579,14 +589,37 @@ def converters( """The conversion functions for the result columns the converter maps, keyed by name.""" description = self.description if self.description else [] return { - d[0]: self._converter.get(d[1]) for d in description if d[1] in self._converter.mappings + name: self._converter.get(d[1]) + for name, d in zip(self._get_column_names(), description, strict=True) + if d[1] in self._converter.mappings } @property def parse_dates(self) -> list[Any | None]: """The names of the result columns with date, time, or timestamp types.""" description = self.description if self.description else [] - return [d[0] for d in description if d[1] in self._PARSE_DATES] + return [ + name + for name, d in zip(self._get_column_names(), description, strict=True) + 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 +674,7 @@ 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]) + return tuple([row[1][name] for name in self._column_names_cache]) def _read_csv(self) -> TextFileReader | DataFrame: import pandas as pd @@ -721,7 +753,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 @@ -957,24 +989,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 7f9bea749..1dbb8e41e 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, @@ -279,7 +281,9 @@ def __init__( 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) + self._df_converters = self._text_value_converters( + self.converters, self._get_column_names() + ) else: self._df = pl.DataFrame() if self._df is not None: @@ -378,8 +382,8 @@ 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 + name: dtype + for name, d in zip(self._get_column_names(), description, strict=True) if (dtype := self._converter.get_dtype(d[1], d[4], d[5])) is not None } @@ -391,16 +395,29 @@ 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 { + name: self._converter.get(d[1]) + for name, d in zip(self._get_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 the 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 _create_dataframe_iterator(self) -> PolarsDataFrameIterator: """Create a DataFrame iterator that reads the result file in chunks. @@ -591,9 +608,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. diff --git a/pyathena/result_set.py b/pyathena/result_set.py index 9ecc91876..81f1167a9 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 f9a667121..42f4ba902 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, json_parse('[2]') 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], + [2], + 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 000000000..c79aa1e29 --- /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 9d0f12794..ec0808139 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -1637,6 +1637,47 @@ 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, 2 AS x, 'a' AS y, json_parse('[1]') AS j, json_parse('[2]') AS j, " + "CAST('12:34:56' AS TIME) AS t, CAST('01:02:03' AS TIME) AS t" + ) + assert pandas_cursor.fetchall() == [ + ( + 1, + 2, + "a", + [1], + [2], + datetime(2017, 1, 1, 12, 34, 56).time(), + datetime(2017, 1, 1, 1, 2, 3).time(), + ) + ] + assert pandas_cursor.as_pandas().columns.tolist() == [ + "x", + "x.1", + "y", + "j", + "j.1", + "t", + "t.1", + ] + @pytest.mark.parametrize( "execute_kwargs", [ diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index a943d2683..9badb479b 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -787,6 +787,47 @@ 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, 2 AS x, 'a' AS y, json_parse('[1]') AS j, json_parse('[2]') AS j, " + "CAST('12:34:56' AS TIME) AS t, CAST('01:02:03' AS TIME) AS t" + ) + assert polars_cursor.fetchall() == [ + ( + 1, + 2, + "a", + [1], + [2], + datetime(2017, 1, 1, 12, 34, 56).time(), + datetime(2017, 1, 1, 1, 2, 3).time(), + ) + ] + assert polars_cursor.as_polars().columns == [ + "x", + "x_duplicated_0", + "y", + "j", + "j_duplicated_0", + "t", + "t_duplicated_0", + ] + @pytest.mark.parametrize( "execute_kwargs", [{}, {"block_size": 2048, "cache_type": "none", "max_workers": 3, "chunksize": 20}], diff --git a/tests/pyathena/test_result_set.py b/tests/pyathena/test_result_set.py index 26635f24d..4eb76777b 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: From 3ee10ab2ae00b91c0c0355c91340559a2450f54f Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 09:10:36 +0900 Subject: [PATCH 2/6] Cover the pyarrow CSV engine fallback for column names that repeat The explicit-engine test builds a result set without __init__, so it now sets _metadata, which the engine check reads through description. Co-Authored-By: Claude Opus 5.5 --- tests/pyathena/pandas/test_cursor.py | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index ec0808139..e488f1869 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -967,6 +967,7 @@ 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 # Test C engine specification result_set._engine = "c" @@ -989,6 +990,22 @@ 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 incompatible chunksize (via parameter) with ( patch.object( From 3216542fdd361a89e0b509671ba0d5945f5137c6 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 10:44:14 +0900 Subject: [PATCH 3/6] Key the CSV column options by the labels that pandas and Polars give the columns - pandas: when the read options parse the result file as Athena writes it, resolve the labels that read_csv() gives the columns, including names=, usecols=, and renamed duplicates, and key dtype, converters, parse_dates, the time truncation, and the binary NULL handling by them. Options given to execute() are kept. Custom quoting and other parsing keeps master's keys, so its values stay text as before. - pandas: fetch rows by position. - Polars: key schema_overrides and the fetch converters by the names of the DataFrame columns, including the new_columns given to execute(). - The public dtypes, converters, and parse_dates properties keep master's keys. Co-Authored-By: Claude Opus 5.5 --- pyathena/pandas/result_set.py | 167 +++++++++++++++++++-------- pyathena/polars/result_set.py | 77 +++++++++--- tests/pyathena/pandas/test_cursor.py | 24 +++- tests/pyathena/polars/test_cursor.py | 16 ++- 4 files changed, 209 insertions(+), 75 deletions(-) diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 3c583d0f5..0bbd8f888 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -411,13 +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] = [ - name - for name, d in zip(self._get_column_names(), description, strict=True) - if d[1] == "time" - ] + self._time_columns: list[Any] = [d[0] for d in description if d[1] == "time"] import pandas as pd @@ -440,9 +437,6 @@ def __init__( # out of the rows that the fetch methods return. Mutable values in its # cells, such as lists from JSON columns, are still shared. self._df_iter = PandasDataFrameIterator(self._df.copy(deep=False), _no_trunc_date) - # Cache column names for fetchone(), after _as_pandas(), which replaces the - # metadata of unload queries. - self._column_names_cache = self._get_column_names() self._iterrows = self._df_iter.iterrows() def _get_parquet_engine(self) -> str: @@ -577,8 +571,8 @@ def dtypes(self) -> dict[str, type[Any]]: """ description = self.description if self.description else [] return { - name: dtype - for name, d in zip(self._get_column_names(), description, strict=True) + d[0]: dtype + for d in description if (dtype := self._converter.get_dtype(d[1], d[4], d[5])) is not None } @@ -589,20 +583,14 @@ def converters( """The conversion functions for the result columns the converter maps, keyed by name.""" description = self.description if self.description else [] return { - name: self._converter.get(d[1]) - for name, d in zip(self._get_column_names(), description, strict=True) - if d[1] in self._converter.mappings + d[0]: self._converter.get(d[1]) for d in description if d[1] in self._converter.mappings } @property def parse_dates(self) -> list[Any | None]: """The names of the result columns with date, time, or timestamp types.""" description = self.description if self.description else [] - return [ - name - for name, d in zip(self._get_column_names(), description, strict=True) - if d[1] in self._PARSE_DATES - ] + 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. @@ -674,7 +662,8 @@ def fetchone( return None else: self._rownumber = row[0] + 1 - return tuple([row[1][name] for name in self._column_names_cache]) + # 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 @@ -708,11 +697,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 @@ -828,12 +820,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 @@ -842,6 +839,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 ( @@ -861,38 +866,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 the column names do not repeat + # (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( diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index 1dbb8e41e..3a89580e1 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -274,15 +274,16 @@ 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. + column_names = self._get_frame_column_names() self._df_converters = self._text_value_converters( - self.converters, self._get_column_names() + self._get_converters(column_names), column_names ) else: self._df = pl.DataFrame() @@ -290,13 +291,13 @@ def __init__( # 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: @@ -381,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 { - name: dtype - for name, d in zip(self._get_column_names(), description, strict=True) - 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]]: @@ -395,13 +392,43 @@ 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 self._get_converters([d[0] for d in description]) + + 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(self._get_column_names(), description, strict=True) + for name, d in zip(column_names, description, strict=True) } def _get_column_names(self) -> list[str]: - """Get the names of the result columns in the DataFrame. + """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``. @@ -419,6 +446,21 @@ def _get_column_names(self) -> list[str]: 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. @@ -437,7 +479,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( @@ -519,7 +562,7 @@ def _read_csv(self) -> pl.DataFrame: lambda: self._csv_storage_options, separator=separator, has_header=has_header, - schema_overrides=self.dtypes, + schema_overrides=self._get_dtypes(self._get_frame_column_names()), ), ) if new_columns: @@ -675,7 +718,7 @@ def _get_csv_params(self) -> tuple[str, bool, list[str] | None]: 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 @@ -711,7 +754,7 @@ def _iter_csv_chunks(self) -> Iterator[pl.DataFrame]: lambda: self._parquet_storage_options, separator=separator, has_header=has_header, - schema_overrides=self.dtypes, + schema_overrides=self._get_dtypes(self._get_frame_column_names()), ), ) for batch in lazy_df.collect_batches(chunk_size=self._chunksize): @@ -779,7 +822,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/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index e488f1869..48e32152a 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -1671,18 +1671,18 @@ def test_fetch_all_rows(self, pandas_cursor): ) def test_duplicate_column_names(self, pandas_cursor): pandas_cursor.execute( - "SELECT 1 AS x, 2 AS x, 'a' AS y, json_parse('[1]') AS j, json_parse('[2]') AS j, " - "CAST('12:34:56' AS TIME) AS t, CAST('01:02:03' AS TIME) AS t" + "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, - 2, "a", + "b", [1], [2], datetime(2017, 1, 1, 12, 34, 56).time(), - datetime(2017, 1, 1, 1, 2, 3).time(), + 2, ) ] assert pandas_cursor.as_pandas().columns.tolist() == [ @@ -1695,6 +1695,22 @@ def test_duplicate_column_names(self, pandas_cursor): "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 9badb479b..68bfe8d14 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -804,18 +804,18 @@ def test_fetch_all_rows(self, polars_cursor): ) def test_duplicate_column_names(self, polars_cursor): polars_cursor.execute( - "SELECT 1 AS x, 2 AS x, 'a' AS y, json_parse('[1]') AS j, json_parse('[2]') AS j, " - "CAST('12:34:56' AS TIME) AS t, CAST('01:02:03' AS TIME) AS t" + "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, - 2, "a", + "b", [1], [2], datetime(2017, 1, 1, 12, 34, 56).time(), - datetime(2017, 1, 1, 1, 2, 3).time(), + 2, ) ] assert polars_cursor.as_polars().columns == [ @@ -828,6 +828,14 @@ def test_duplicate_column_names(self, polars_cursor): "t_duplicated_0", ] + 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}], From 0e1cc83024adedb2b833c973297182eea1755130 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 11:18:24 +0900 Subject: [PATCH 4/6] Rename Polars columns after reading, and convert Arrow rows by position - Polars: key schema_overrides by the CSV header and apply the new_columns given to execute() after reading. read_csv() and scan_csv() interpret a schema_overrides dict differently when new_columns renames only the first columns. - Arrow: choose the fetch converters by column position, so that columns with the same name and different types, such as json and varchar, keep their own conversion. - pandas: do not use the pyarrow engine when the column labels need resolving, because it applied the types by the original names to the columns that names= renamed. Co-Authored-By: Claude Opus 5.5 --- pyathena/arrow/result_set.py | 24 +++++------ pyathena/pandas/result_set.py | 13 +++--- pyathena/polars/result_set.py | 55 +++++++++++++----------- tests/pyathena/arrow/test_cursor.py | 4 +- tests/pyathena/pandas/test_cursor.py | 13 ++++++ tests/pyathena/polars/test_cursor.py | 6 +++ tests/pyathena/polars/test_result_set.py | 6 +-- 7 files changed, 75 insertions(+), 46 deletions(-) diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index 12d82e671..8fb6c041b 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,19 +254,19 @@ def _fetch(self) -> None: except StopIteration: return else: - # Read the columns by position; to_pydict() keeps one column per name. + # 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] - converters = ( - self.converters - if self._convert_rows - else self._text_value_converters(self.converters) - ) - if converters: - column_converters = [ - converters.get(name, _to_default) for name in rows.schema.names - ] + 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(convert(v) for convert, v in zip(column_converters, row, strict=False)) + tuple(convert(v) for convert, v in zip(converters, row, strict=False)) for row in zip(*columns, strict=False) ] else: diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 0bbd8f888..2112773b7 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -472,13 +472,16 @@ def _get_csv_engine( # checks pass; otherwise fall through to the C engine default. if self._engine == "pyarrow": effective_chunksize = chunksize if chunksize is not None else self._chunksize - column_names = [d[0] for d in self.description or []] is_compatible = ( effective_chunksize is None and self._quoting == 1 and not self.converters - # The pyarrow engine does not rename columns with the same name. - and len(set(column_names)) == len(column_names) + # 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: @@ -883,8 +886,8 @@ def _get_csv_column_labels( """ import pandas as pd - # The pyarrow engine runs only when the column names do not repeat - # (see _get_csv_engine()), and does not support reading only the header. + # 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 []] diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index 3a89580e1..aadcb0e8b 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -394,6 +394,11 @@ def converters(self) -> dict[str, Callable[[str | None], Any | None]]: description = self.description if self.description else [] 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. @@ -554,19 +559,19 @@ 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, + ) + # 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._get_dtypes(self._get_frame_column_names()), - ), - ) + 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}.") @@ -713,7 +718,9 @@ 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()``. """ if self.output_location and self.output_location.endswith(".txt"): separator = "\t" @@ -722,7 +729,7 @@ def _get_csv_params(self) -> tuple[str, bool, list[str] | None]: else: separator = "," has_header = True - new_columns = None + new_columns = self._kwargs.get("new_columns") return separator, has_header, new_columns def _iter_csv_chunks(self) -> Iterator[pl.DataFrame]: @@ -744,22 +751,22 @@ 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, + ) + # 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._get_dtypes(self._get_frame_column_names()), - ), - ) + 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}.") diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index 42f4ba902..be738f3a2 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -1071,7 +1071,7 @@ def test_fetch_all_rows(self, 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, json_parse('[2]') AS j, " + "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() == [ @@ -1080,7 +1080,7 @@ def test_duplicate_column_names(self, arrow_cursor): 2, "a", [1], - [2], + "b", datetime(2017, 1, 1, 12, 34, 56).time(), datetime(2017, 1, 1, 1, 2, 3).time(), ) diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 48e32152a..96fb6f238 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -968,6 +968,7 @@ def test_get_csv_engine_explicit_specification(self): 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" @@ -1006,6 +1007,18 @@ def test_get_csv_engine_explicit_specification(self): 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( diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index 68bfe8d14..2dac0b525 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -828,6 +828,12 @@ def test_duplicate_column_names(self, polars_cursor): "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_duplicate_column_names_new_columns(self, polars_cursor): """The column types follow the columns that new_columns renames.""" polars_cursor.execute( diff --git a/tests/pyathena/polars/test_result_set.py b/tests/pyathena/polars/test_result_set.py index c16318ecf..946fabafc 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}, ), @@ -146,7 +146,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, From fb7bfab15c214d59c2a6170f8bb3fd1d8d77405b Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 11:32:43 +0900 Subject: [PATCH 5/6] Let Polars rename the columns when execute() replaces the schema overrides schema_overrides given to execute() replace the result set's types and may be keyed by the new_columns names, which Polars resolves itself. Co-Authored-By: Claude Opus 5.5 --- pyathena/polars/result_set.py | 18 ++++++++++++------ tests/pyathena/polars/test_cursor.py | 7 +++++++ 2 files changed, 19 insertions(+), 6 deletions(-) diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index aadcb0e8b..c2ef562c7 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -565,8 +565,9 @@ def _read_csv(self) -> pl.DataFrame: has_header=has_header, schema_overrides=self._csv_dtypes, ) - # Renamed after reading, so that Polars matches the types to the header. - read_kwargs.pop("new_columns", None) + if new_columns: + # 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, **read_kwargs) @@ -720,7 +721,9 @@ def _get_csv_params(self) -> tuple[str, bool, list[str] | None]: Returns: 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()``. + 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" @@ -729,7 +732,9 @@ def _get_csv_params(self) -> tuple[str, bool, list[str] | None]: else: separator = "," has_header = True - new_columns = self._kwargs.get("new_columns") + 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]: @@ -759,8 +764,9 @@ def _iter_csv_chunks(self) -> Iterator[pl.DataFrame]: has_header=has_header, schema_overrides=self._csv_dtypes, ) - # Renamed after reading, so that Polars matches the types to the header. - read_kwargs.pop("new_columns", None) + if new_columns: + # Renamed after reading, so that Polars matches the types to the header. + read_kwargs.pop("new_columns", None) try: lazy_df = pl.scan_csv(self.output_location, **read_kwargs) diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index 2dac0b525..d4c642eb6 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -834,6 +834,13 @@ def test_new_columns_renaming_first_columns(self, polars_cursor, chunksize): 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( From e78322210ae9eee0112e742ed60e309aecaf437d Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sun, 4 Oct 2026 11:41:03 +0900 Subject: [PATCH 6/6] Leave new_columns to Polars for .txt results when execute() gives the types Co-Authored-By: Claude Opus 5.5 --- pyathena/polars/result_set.py | 4 +-- tests/pyathena/polars/test_result_set.py | 36 ++++++++++++++++++++++++ 2 files changed, 38 insertions(+), 2 deletions(-) diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index c2ef562c7..85ea4742b 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -565,7 +565,7 @@ def _read_csv(self) -> pl.DataFrame: has_header=has_header, schema_overrides=self._csv_dtypes, ) - if new_columns: + 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) @@ -764,7 +764,7 @@ def _iter_csv_chunks(self) -> Iterator[pl.DataFrame]: has_header=has_header, schema_overrides=self._csv_dtypes, ) - if new_columns: + 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) diff --git a/tests/pyathena/polars/test_result_set.py b/tests/pyathena/polars/test_result_set.py index 946fabafc..231f8daa8 100644 --- a/tests/pyathena/polars/test_result_set.py +++ b/tests/pyathena/polars/test_result_set.py @@ -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"), [