diff --git a/docs/pandas.md b/docs/pandas.md index bdb0f749..1ceb6d7e 100644 --- a/docs/pandas.md +++ b/docs/pandas.md @@ -556,8 +556,9 @@ Common performance options: With `engine="pyarrow"`, PandasCursor uses the PyArrow engine only when pyarrow is installed, no chunksize is set (explicitly or by `auto_optimize_chunksize`), `quoting` is the default, the result has no columns that need a converter (`boolean`, `decimal`, `varbinary`, `json`, `time with time zone`, and `timestamp with time zone` with the default converter), and the result file is at least `AthenaPandasResultSet.PYARROW_MIN_FILE_SIZE_BYTES` bytes. Otherwise, it falls back to the C engine. -With the PyArrow engine, when `keep_default_na` and `na_values` are the defaults and the only pandas.read_csv() options passed to `execute()` are `dtype` as a mapping and `parse_dates` as a list, PyAthena reads the file with `pyarrow.csv`, allowing newlines in quoted values, and converts it to a DataFrame as `pandas.read_csv(engine="pyarrow")` does. +With the PyArrow engine, when `keep_default_na` and `na_values` are the defaults and the only pandas.read_csv() options passed to `execute()` are `dtype` as a mapping and `parse_dates` as a list, PyAthena reads the file with `pyarrow.csv`, allowing newlines in quoted values, and converts it to a DataFrame as `pandas.read_csv(engine="pyarrow")` does, except that columns with a string dtype are read as text, other than in the tab-separated results of DDL statements, and without `future.infer_string`, a column with dtype `str` keeps NULL as a missing value. Otherwise, pandas reads the file, and its PyArrow engine raises an error or returns wrong values when a quoted value containing a newline crosses one of pyarrow's read blocks. +It also changes the values of a column with a string dtype whose values all look like numbers, such as `"007"` to `"7.0"`, and without `future.infer_string`, turns NULL in a column with dtype `str` into the string `"nan"`. Apart from PandasCursor's own options such as `engine` and `chunksize`, an option passed here replaces the value PyAthena sets for the same pandas.read_csv() argument. For example, `dtype` replaces the whole column type mapping, and `parse_dates` replaces the list of date, time, and timestamp columns. diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 2ae91078..3f21f8a5 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -56,6 +56,10 @@ def _read_csv_with_pyarrow(source: str | IOBase, read_csv_kwargs: dict[str, Any] ``parse_dates`` columns are parsed with ``pandas.to_datetime()``, and without ``future.infer_string``, strings become objects. + Unlike pandas' PyArrow engine, columns with a string dtype are read as text, + except in a header-less file, and a NumPy str dtype keeps missing values + missing, as with pandas' C engine. + Args: source: The result file. read_csv_kwargs: The pandas.read_csv() options built by @@ -71,6 +75,25 @@ def _read_csv_with_pyarrow(source: str | IOBase, read_csv_kwargs: dict[str, Any] header = read_csv_kwargs["header"] names = read_csv_kwargs["names"] null_values = list(read_csv_kwargs["na_values"]) + string_columns = set() + for column, value in read_csv_kwargs["dtype"].items(): + try: + column_dtype = pd.api.types.pandas_dtype(value) + except (TypeError, ValueError, NotImplementedError): + # pandas validates only the entries of columns in the result. + continue + if isinstance(column_dtype, pd.StringDtype) or column_dtype.kind == "U": + string_columns.add(column) + # Columns with a string dtype keep their text instead of the type pyarrow infers, + # so "007" stays "007" as with pandas' C engine. The fields of a header-less file + # get their names only after reading, as pandas gives the names to the last + # fields, so they keep the inferred types. pandas ignores dtype keys that are not + # column names, such as positions. + column_types = ( + {} + if header is None + else {column: pa.string() for column in string_columns if isinstance(column, str)} + ) table = pyarrow_csv.read_csv( source, read_options=pyarrow_csv.ReadOptions(autogenerate_column_names=header is None), @@ -80,7 +103,9 @@ def _read_csv_with_pyarrow(source: str | IOBase, read_csv_kwargs: dict[str, Any] newlines_in_values=True, ), convert_options=pyarrow_csv.ConvertOptions( - null_values=null_values, strings_can_be_null="" in null_values + column_types=column_types, + null_values=null_values, + strings_can_be_null="" in null_values, ), ) schema = table.schema @@ -109,6 +134,15 @@ def _read_csv_with_pyarrow(source: str | IOBase, read_csv_kwargs: dict[str, Any] for column, value in dtype.items() if column in df.columns } + # astype() with a NumPy str dtype, which str means without future.infer_string, + # turns missing values into strings such as "nan", so these columns are + # converted after the dates instead. + numpy_string_columns = { + column + for column, value in dtype.items() + if not isinstance(value, pd.api.extensions.ExtensionDtype) and value.kind == "U" + } + dtype = {column: value for column, value in dtype.items() if column not in numpy_string_columns} df = df.astype(dtype) if not pd.get_option("future.infer_string"): # Without the string dtype, pandas returns strings, and the string @@ -138,6 +172,12 @@ def _read_csv_with_pyarrow(source: str | IOBase, read_csv_kwargs: dict[str, Any] df[column] = values.to_numpy(dtype=object, na_value=float("nan")) # pandas applies the dtype mapping again after parsing dates. df = df.astype(dtype) + for index, column in enumerate(df.columns): + if column in numpy_string_columns: + # Object strings with missing values kept, as pandas' C engine returns. + values = df.iloc[:, index] + strings = values.astype(str).astype(object) + df.isetitem(index, strings.where(values.notna(), float("nan"))) return df diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 44bbf259..26e70154 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -327,6 +327,32 @@ def test_csv_storage_options(self, pandas_cursor, query, expected, binary, with_ # pandas opens and closes the file itself. assert pandas_cursor.result_set._csv_stream is None + @pytest.mark.parametrize("infer_string", [True, False]) + def test_pyarrow_engine_string_values(self, pandas_cursor, infer_string): + # The PyArrow engine keeps the text of strings that look like numbers, and + # a NULL string missing, as the C engine does. + query = """ + SELECT * FROM (VALUES + (1, '1', 'abcdefghijklmnopqrstuvwxyz'), + (2, CAST(NULL AS VARCHAR), 'abcdefghijklmnopqrstuvwxyz'), + (3, 'nan', 'abcdefghijklmnopqrstuvwxyz'), + (4, '007', 'abcdefghijklmnopqrstuvwxyz'), + (5, '1e3', 'abcdefghijklmnopqrstuvwxyz') + ) AS t(id, v, padding) ORDER BY id + """ + with pd.option_context("future.infer_string", infer_string): + with patch( + "pyathena.pandas.result_set._read_csv_with_pyarrow", wraps=_read_csv_with_pyarrow + ) as read_csv_with_pyarrow: + pandas_cursor.execute(query, engine="pyarrow") + actual = pandas_cursor.as_pandas()["v"] + read_csv_with_pyarrow.assert_called_once() + pandas_cursor.execute(query, engine="c") + expected = pandas_cursor.as_pandas()["v"] + pd.testing.assert_series_equal(actual, expected, check_exact=True) + assert actual.iloc[[0, 2, 3, 4]].tolist() == ["1", "nan", "007", "1e3"] + assert pd.isna(actual.iloc[1]) + def test_pyarrow_engine_multiline_values_across_blocks(self, pandas_cursor): # The 2.4 MB result spans several 1 MiB pyarrow read blocks, and its # values contain a newline and quotes. diff --git a/tests/pyathena/pandas/test_result_set.py b/tests/pyathena/pandas/test_result_set.py index ef6022a5..0ed38d81 100644 --- a/tests/pyathena/pandas/test_result_set.py +++ b/tests/pyathena/pandas/test_result_set.py @@ -10,6 +10,7 @@ from unittest.mock import MagicMock, PropertyMock, patch import pandas as pd +import pyarrow as pa import pytest from pandas.testing import assert_frame_equal @@ -153,6 +154,15 @@ def test_read_parquet_filesystem(self, execute_kwargs, path, filesystem_kwargs): } +def _is_string_dtype(value): + """Return whether a dtype mapping value is a string dtype, ignoring invalid values.""" + try: + dtype = pd.api.types.pandas_dtype(value) + except TypeError: + return False + return isinstance(dtype, pd.StringDtype) or dtype.kind == "U" + + def _pyarrow_read_csv_kwargs(types, tab_separated=False, **kwargs): """Build the pandas.read_csv() options with AthenaPandasResultSet._get_csv_read_options(). @@ -224,6 +234,32 @@ def _pyarrow_read_csv_kwargs(types, tab_separated=False, **kwargs): '"x","x","d"\n"1","2","2024-01-01"\n,,\n', _pyarrow_read_csv_kwargs({"x": "integer", "d": "date"}), ), + ( + '"v","n"\n"1","1"\n,\n"nan","3"\n"007","4"\n"1e3","5"\n', + _pyarrow_read_csv_kwargs({"v": "varchar", "n": "integer"}), + ), + ( + '"v","w","x"\n"007","a","1"\n,,"2"\n', + _pyarrow_read_csv_kwargs( + {"x": "integer"}, + dtype={ + "v": pd.ArrowDtype(pa.string()), + "w": pd.ArrowDtype(pa.large_string()), + "x": pd.Int64Dtype(), + }, + ), + ), + ( + "001\t2\t003\n004\t5\t\n", + _pyarrow_read_csv_kwargs({"v": "varchar"}, True), + ), + ( + '"v","n"\n"007","1"\n', + _pyarrow_read_csv_kwargs( + {"v": "varchar", "n": "integer"}, + dtype={**_pyarrow_read_csv_kwargs({"v": "varchar"})["dtype"], 0: str}, + ), + ), ( '"v"\n"plain"\n"2024-01-01"\n\n', _pyarrow_read_csv_kwargs({"v": "varchar"}, parse_dates=["v"]), @@ -244,6 +280,10 @@ def _pyarrow_read_csv_kwargs(types, tab_separated=False, **kwargs): "dtype_of_date_column", "dtype_none", "duplicate_names", + "numeric_looking_strings", + "arrow_string_dtypes", + "tab_separated_numeric_fields", + "dtype_position_key", "unparsed_dates", "tab_separated_extra_fields", "tab_separated", @@ -251,12 +291,32 @@ def _pyarrow_read_csv_kwargs(types, tab_separated=False, **kwargs): ) def test_read_csv_with_pyarrow_matches_pandas(data, read_csv_kwargs, infer_string): # Without values that cross a read block, the result is the one of - # pandas.read_csv(engine="pyarrow"). + # pandas.read_csv(engine="pyarrow"), except that the columns with a string + # dtype have the values of pandas' C engine, and in a header-less file, its + # missing values. Where the C engine parses a parse_dates column despite its + # string dtype, the PyArrow engine's applying the dtype again is kept. with pd.option_context("future.infer_string", infer_string): expected = pd.read_csv( io.BytesIO(data.encode()), **{**read_csv_kwargs, "dtype": dict(read_csv_kwargs["dtype"])}, ) + c_engine = pd.read_csv( + io.BytesIO(data.encode()), + **{**read_csv_kwargs, "engine": "c", "dtype": dict(read_csv_kwargs["dtype"])}, + ) + string_columns = { + column for column, value in read_csv_kwargs["dtype"].items() if _is_string_dtype(value) + } + for index, column in enumerate(expected.columns): + if column not in string_columns or c_engine[column].dtype.kind == "M": + continue + if read_csv_kwargs["header"] is None: + # Header-less fields keep the inferred types, and only their missing + # values follow the C engine, which makes extra fields the index. + missing = c_engine[column].isna().to_numpy() + expected.isetitem(index, expected.iloc[:, index].mask(missing, float("nan"))) + else: + expected.isetitem(index, c_engine[column].array) actual = _read_csv_with_pyarrow( io.BytesIO(data.encode()), {**read_csv_kwargs, "dtype": dict(read_csv_kwargs["dtype"])}, @@ -264,6 +324,35 @@ def test_read_csv_with_pyarrow_matches_pandas(data, read_csv_kwargs, infer_strin assert_frame_equal(actual, expected, check_exact=True) +@pytest.mark.parametrize("infer_string", [True, False]) +def test_read_csv_with_pyarrow_string_dtype_after_dates(infer_string): + # A string dtype applies again after parse_dates, as with pandas' PyArrow + # engine, and keeps NULL missing. + with pd.option_context("future.infer_string", infer_string): + df = _read_csv_with_pyarrow( + io.BytesIO(b'"v"\n"2024-01-01"\n\n'), + _pyarrow_read_csv_kwargs({"v": "varchar"}, parse_dates=["v"]), + ) + assert df["v"].tolist()[0] == "2024-01-01" + assert pd.isna(df["v"].tolist()[1]) + + +def test_read_csv_with_pyarrow_ignores_unused_dtype_entries(): + # As with pandas' PyArrow engine, a dtype entry for a column that is not in + # the result is not validated. + read_csv_kwargs = _pyarrow_read_csv_kwargs( + {"v": "varchar"}, + dtype={"v": str, "unused": "not-a-dtype", "unsupported": "decimal128(10, 2)[pyarrow]"}, + ) + data = b'"v"\n"007"\n' + expected = pd.read_csv(io.BytesIO(data), **{**read_csv_kwargs, "dtype": {"v": str}}) + actual = _read_csv_with_pyarrow( + io.BytesIO(data), {**read_csv_kwargs, "dtype": dict(read_csv_kwargs["dtype"])} + ) + assert actual.columns.tolist() == expected.columns.tolist() == ["v"] + assert actual["v"].tolist() == ["007"] + + def test_read_csv_with_pyarrow_multiline_values_across_blocks(): # The 2.4 MB file spans several 1 MiB pyarrow read blocks, and its values # contain newlines and quotes.