diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index 8fb6c041..d34a7cb5 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -302,10 +302,24 @@ def _read_csv(self) -> Table: ): return pa.Table.from_pydict({}) length = self._get_content_length() - binary_columns = {d[0] for d in self.description or [] if d[1] == "varbinary"} + description = self.description if self.description else [] + names = [d[0] for d in description] + # pyarrow types every column with a name by its column_types entry, so columns + # with the same name are read under their positions and get their names back + # after reading. + has_duplicate_names = len(set(names)) != len(names) + if has_duplicate_names: + column_names = [str(i) for i in range(len(names))] + column_types = { + str(i): dtype + for i, d in enumerate(description) + if (dtype := self._converter.get_dtype(d[1], d[4], d[5])) is not None + } + else: + column_names = names + column_types = self.column_types + binary_columns = {i for i, d in enumerate(description) if d[1] == "varbinary"} if length and self.output_location.endswith(".txt"): - description = self.description if self.description else [] - column_names = [d[0] for d in description] read_opts = csv.ReadOptions( skip_rows=0, column_names=column_names, @@ -320,6 +334,11 @@ def _read_csv(self) -> Table: ) elif length and self.output_location.endswith(".csv"): read_opts = csv.ReadOptions(skip_rows=0, block_size=self._block_size, use_threads=True) + if has_duplicate_names: + read_opts.column_names = column_names + # Skips the header as a parsed row; skip_rows would split a quoted name + # that contains a newline. + read_opts.skip_rows_after_names = 1 parse_opts = csv.ParseOptions( delimiter=",", quote_char='"', @@ -344,12 +363,14 @@ def _read_csv(self) -> Table: strings_can_be_null=bool(binary_columns), quoted_strings_can_be_null=False, timestamp_parsers=self.timestamp_parsers, - column_types=self.column_types, + column_types=column_types, ), ) + if has_duplicate_names: + table = table.rename_columns(names) if binary_columns: for index, field in enumerate(table.schema): - if field.name not in binary_columns and ( + if index not in binary_columns and ( pa.types.is_string(field.type) or pa.types.is_binary(field.type) ): # Preserve the existing CSV behavior for non-binary Athena columns. diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index c47780b0..2ae91078 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -429,6 +429,32 @@ class AthenaPandasResultSet(AthenaResultSet): ] # The pandas.read_csv() options given to execute() that _read_csv_with_pyarrow() reads. _PYARROW_READ_CSV_OPTIONS: ClassVar[frozenset[str]] = frozenset({"dtype", "parse_dates"}) + # The pandas.read_csv() options given to execute() that do not change how pandas reads + # the header row, with which _read_csv_header_as_labels() replaces the header. + _LABELED_HEADER_READ_CSV_OPTIONS: ClassVar[frozenset[str]] = frozenset( + { + "cache_dates", + "converters", + "date_format", + "dayfirst", + "decimal", + "dtype_backend", + "false_values", + "float_precision", + "index_col", + "keep_default_na", + "low_memory", + "na_filter", + "na_values", + "nrows", + "on_bad_lines", + "parse_dates", + "skipfooter", + "thousands", + "true_values", + "usecols", + } + ) def __init__( self, @@ -808,6 +834,9 @@ def _read_csv(self) -> TextFileReader | DataFrame: with ExitStack() as stack: source: str | IOBase = self.output_location binary_columns = self._configure_binary_csv_read(read_csv_kwargs, labels) + if labels is not None: + # After _configure_binary_csv_read(), which checks for the header row. + self._read_csv_header_as_labels(read_csv_kwargs, csv_engine) self._csv_converters = read_csv_kwargs.get("converters") or {} if binary_columns: # Given storage_options, even None, open the file through fsspec @@ -1058,6 +1087,41 @@ def _key_csv_columns_by_labels( ] self._time_columns = [label for label, d in columns if d[1] == "time"] + def _read_csv_header_as_labels(self, read_csv_kwargs: dict[str, Any], csv_engine: str) -> None: + """Read the header of the CSV result file as the labels of columns with the same name. + + pandas gives a column that it renames, such as ``x.1``, the dtype of the first + column with the name when it has none of its own. When PyAthena builds the + dtypes, the columns are read under their labels, so that each keeps its own type. + + Args: + read_csv_kwargs: The options for ``pandas.read_csv()``, updated in place. + csv_engine: The CSV engine that reads the file. + """ + import pandas as pd + + names = [d[0] for d in self.description or []] + if len(set(names)) == len(names): + return + if not self._kwargs.keys() <= self._LABELED_HEADER_READ_CSV_OPTIONS: + # Other options, such as dtype, names, sep, comment, or encoding, keep the + # header row as pandas reads it. + return + # pandas renames the names in a header row, and copies their dtypes, even with + # names given, so the header row is skipped instead. + read_csv_kwargs["names"] = self._resolve_csv_column_names( + names, read_csv_kwargs, pd.read_csv + )[0] + read_csv_kwargs["header"] = None + if csv_engine == "python": + # The python engine skips lines, which a name with a newline spans. + header = StringIO(newline="") + csv.writer(header, quoting=csv.QUOTE_ALL, lineterminator="").writerow(names) + read_csv_kwargs["skiprows"] = len(StringIO(header.getvalue(), newline="").readlines()) + else: + # The C engine skips parsed rows. + read_csv_kwargs["skiprows"] = 1 + def _configure_binary_csv_read( self, read_csv_kwargs: dict[str, Any], labels: list[Any] | None ) -> set[int]: diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index be738f3a..54bc9253 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -32,13 +32,14 @@ class TestArrowCursor: def test_binary_null_vs_empty(self, arrow_cursor): - query = """SELECT * FROM (VALUES + # The text column has the name of the binary column, and keeps its own NULL. + query = """SELECT id, value, label, text_value AS value FROM (VALUES (1, CAST(NULL AS VARBINARY), 'null', CAST(NULL AS VARCHAR)), (2, X'', 'empty', ''), (3, X'00ff275c25', 'comma, quote" and' || chr(10) || 'newline', 'NULL') ) AS t(id, value, label, text_value) ORDER BY id""" arrow_cursor.execute(query) - assert arrow_cursor.as_arrow().column("value").to_pylist() == [None, "", "00 ff 27 5c 25"] + assert arrow_cursor.as_arrow().column(1).to_pylist() == [None, "", "00 ff 27 5c 25"] rows = arrow_cursor.fetchall() assert [row[:3] for row in rows] == [ (1, None, "null"), @@ -1071,13 +1072,15 @@ 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, 'b' AS j, " - "CAST('12:34:56' AS TIME) AS t, CAST('01:02:03' AS TIME) AS t" + "SELECT 1 AS x, 2 AS x, CAST('01:02:03' AS TIME) 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, + datetime(2017, 1, 1, 1, 2, 3).time(), "a", [1], "b", @@ -1085,7 +1088,13 @@ def test_duplicate_column_names(self, arrow_cursor): datetime(2017, 1, 1, 1, 2, 3).time(), ) ] - assert arrow_cursor.as_arrow().column_names == ["x", "x", "y", "j", "j", "t", "t"] + assert arrow_cursor.as_arrow().column_names == ["x", "x", "x", "y", "j", "j", "t", "t"] + + def test_duplicate_column_names_with_newline(self, arrow_cursor): + """Columns with the same name keep their own types, with a newline in the name.""" + arrow_cursor.execute('SELECT 1 AS "a\nb", CAST(\'01:02:03\' AS TIME) AS "a\nb"') + assert arrow_cursor.fetchall() == [(1, datetime(2017, 1, 1, 1, 2, 3).time())] + assert arrow_cursor.as_arrow().column_names == ["a\nb", "a\nb"] @pytest.mark.parametrize( "execute_kwargs", [{}, {"connect_timeout": 3.0, "request_timeout": 4.0}] diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 5a9b9424..00bffbfa 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -1743,13 +1743,14 @@ def test_fetch_all_rows(self, 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" + "SELECT 1 AS x, 'a' AS x, CAST('01:02:03' AS TIME) 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", + datetime(2017, 1, 1, 1, 2, 3).time(), "b", [1], [2], @@ -1760,6 +1761,7 @@ def test_duplicate_column_names(self, pandas_cursor): assert pandas_cursor.as_pandas().columns.tolist() == [ "x", "x.1", + "x.2", "y", "j", "j.1", @@ -1767,6 +1769,45 @@ def test_duplicate_column_names(self, pandas_cursor): "t.1", ] + @pytest.mark.parametrize("engine", ["c", "python"]) + def test_duplicate_column_names_with_newline(self, pandas_cursor, engine): + """Columns with the same name keep their own types, with a newline in the name.""" + pandas_cursor.execute( + 'SELECT 1 AS "a\nb", CAST(\'01:02:03\' AS TIME) AS "a\nb", ' + "INTERVAL '2' DAY AS \"a\nb\"", + engine=engine, + ) + assert pandas_cursor.fetchall() == [ + (1, datetime(2017, 1, 1, 1, 2, 3).time(), "2 00:00:00.000") + ] + assert pandas_cursor.as_pandas().columns.tolist() == ["a\nb", "a\nb.1", "a\nb.2"] + + @pytest.mark.parametrize( + ("query", "execute_kwargs", "expected_rows", "expected_columns"), + [ + ( + 'SELECT CAST(\'01:02:03\' AS TIME) AS "a\\b", 1 AS "a\\b"', + {"engine": "python", "escapechar": "\\"}, + [(datetime(2017, 1, 1, 1, 2, 3).time(), 1)], + ["ab", "ab.1"], + ), + ( + "SELECT 1 AS x, 2 AS x WHERE false", + {"engine": "python", "sep": None}, + [], + ["x", "x.1"], + ), + ], + ids=["escapechar", "detected_delimiter"], + ) + def test_duplicate_column_names_header_options( + self, pandas_cursor, query, execute_kwargs, expected_rows, expected_columns + ): + """Options that change how pandas reads the header still apply to the column labels.""" + pandas_cursor.execute(query, **execute_kwargs) + assert pandas_cursor.fetchall() == expected_rows + assert pandas_cursor.as_pandas().columns.tolist() == expected_columns + @pytest.mark.parametrize( ("execute_kwargs", "expected_row", "expected_columns"), [