diff --git a/docs/pandas.md b/docs/pandas.md index 20387867..bdb0f749 100644 --- a/docs/pandas.md +++ b/docs/pandas.md @@ -556,6 +556,8 @@ 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. +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. 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/docs/polars.md b/docs/polars.md index 9fcc692a..f347e421 100644 --- a/docs/polars.md +++ b/docs/polars.md @@ -15,6 +15,11 @@ does not require PyArrow as a dependency. CSV results read without the chunksize PyAthena's own S3FileSystem (fsspec compatible), and other reads use Polars' native S3 access, so s3fs is also not required. +Options passed to `execute()`, other than PolarsCursor's own options such as `chunksize`, are passed to the Polars read function. +`pl.read_csv()` reads the CSV result with pyarrow when `use_pyarrow=True` is given with `schema_overrides=None`, which replaces PyAthena's column types, and Polars' other conditions for it hold. +That read raises an error or returns wrong values when a quoted value containing a newline crosses one of pyarrow's read blocks. +The default Polars CSV reader reads such values. + You can use the PolarsCursor by specifying the `cursor_class` with the connect method or connection object. diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index d700fa72..3b87348a 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -43,6 +43,104 @@ def _no_trunc_date(df: DataFrame) -> DataFrame: return df +def _read_csv_with_pyarrow(source: str | IOBase, read_csv_kwargs: dict[str, Any]) -> DataFrame: + """Read a CSV result with pyarrow as ``pandas.read_csv(engine="pyarrow")`` does. + + pandas' PyArrow engine does not set ``newlines_in_values``, so pyarrow + raises or returns wrong values when a quoted value containing a newline + crosses one of its read blocks. This reads the file with that option and + converts the table as pandas does for the options that + ``AthenaPandasResultSet._reads_csv_with_pyarrow()`` accepts: NULL-typed + columns become float64, integer columns without a ``dtype`` entry get + NumPy integer types, the ``dtype`` mapping is applied before and after the + ``parse_dates`` columns are parsed with ``pandas.to_datetime()``, and + without ``future.infer_string``, strings become objects. + + Args: + source: The result file. + read_csv_kwargs: The pandas.read_csv() options built by + ``AthenaPandasResultSet._get_csv_read_options()``. + + Returns: + The result as a DataFrame. + """ + import pandas as pd + import pyarrow as pa + from pyarrow import csv as pyarrow_csv + + header = read_csv_kwargs["header"] + names = read_csv_kwargs["names"] + null_values = list(read_csv_kwargs["na_values"]) + table = pyarrow_csv.read_csv( + source, + read_options=pyarrow_csv.ReadOptions(autogenerate_column_names=header is None), + parse_options=pyarrow_csv.ParseOptions( + delimiter=read_csv_kwargs["sep"], + ignore_empty_lines=read_csv_kwargs["skip_blank_lines"], + newlines_in_values=True, + ), + convert_options=pyarrow_csv.ConvertOptions( + null_values=null_values, strings_can_be_null="" in null_values + ), + ) + schema = table.schema + for index, type_ in enumerate(schema.types): + if pa.types.is_null(type_): + schema = schema.set(index, schema.field(index).with_type(pa.float64())) + integer_dtypes = { + pa.int8(): pd.Int8Dtype(), + pa.int16(): pd.Int16Dtype(), + pa.int32(): pd.Int32Dtype(), + pa.int64(): pd.Int64Dtype(), + } + # Integers become nullable dtypes first so that a dtype entry converts them + # without going through float64. + df = table.cast(schema).to_pandas(types_mapper=integer_dtypes.get) + if header is None: + # pandas names the columns beyond the given names by their positions. + df.columns = [str(index) for index in range(len(df.columns) - len(names))] + names + dtype = dict(read_csv_kwargs["dtype"]) + for column in df.columns: + # Integer columns without a dtype entry get NumPy integer types. + if column not in dtype and df[column].dtype in integer_dtypes.values(): + dtype[column] = df[column].dtype.numpy_dtype + dtype = { + column: pd.api.types.pandas_dtype(value) + for column, value in dtype.items() + if column in df.columns + } + df = df.astype(dtype) + if not pd.get_option("future.infer_string"): + # Without the string dtype, pandas returns strings, and the string + # categories of categorical columns, as objects. + for index in range(len(df.columns)): + values = df.iloc[:, index] + if values.dtype == "str": + df.isetitem(index, values.astype(object).fillna(None)) + elif isinstance(values.dtype, pd.CategoricalDtype) and ( + values.dtype.categories.dtype == "str" + ): + categories = values.dtype.categories.astype(object) + df.isetitem( + index, + values.astype(pd.CategoricalDtype(categories, ordered=values.dtype.ordered)), + ) + for column in read_csv_kwargs["parse_dates"]: + if isinstance(column, int) and column not in df.columns: + column = df.columns[column] + if df[column].dtype.kind in "Mm": + continue + values = df[column].astype("string") + try: + df[column] = pd.to_datetime(values, utc=False) + except (ValueError, TypeError): + # pandas keeps the column as strings if it cannot parse it. + df[column] = values.to_numpy(dtype=object, na_value=float("nan")) + # pandas applies the dtype mapping again after parsing dates. + df = df.astype(dtype) + return df + + class _JSONConverter: """A json converter for ``pandas.read_csv()`` that keeps NULL from making values floats. @@ -329,6 +427,8 @@ class AthenaPandasResultSet(AthenaResultSet): "time", "timestamp", ] + # 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"}) def __init__( self, @@ -696,7 +796,10 @@ def _read_csv(self) -> TextFileReader | DataFrame: source = self._csv_stream = stack.enter_context( self._fs.open(self.output_location, mode="rb") ) - result = pd.read_csv(source, **read_csv_kwargs) + if csv_engine == "pyarrow" and self._reads_csv_with_pyarrow(): + result = _read_csv_with_pyarrow(source, read_csv_kwargs) + else: + result = pd.read_csv(source, **read_csv_kwargs) if not isinstance(result, pd.DataFrame): # The chunk iterator takes ownership of the stream. stack.pop_all() @@ -716,6 +819,25 @@ def _read_csv(self) -> TextFileReader | DataFrame: _logger.exception(f"Failed to read {self.output_location}.") raise OperationalError(*e.args) from e + def _reads_csv_with_pyarrow(self) -> bool: + """Whether ``_read_csv_with_pyarrow()`` reads the CSV result for the PyArrow engine. + + It reproduces ``pandas.read_csv(engine="pyarrow")`` for PyAthena's default NA + values and for ``dtype`` as a mapping and ``parse_dates`` as a list given to + ``execute()``. With other options, pandas reads the file. + + Returns: + True if ``_read_csv_with_pyarrow()`` reads the result. + """ + return ( + not self._keep_default_na + and isinstance(self._na_values, (list, tuple)) + and list(self._na_values) == [""] + and self._kwargs.keys() <= self._PYARROW_READ_CSV_OPTIONS + and isinstance(self._kwargs.get("dtype", {}), dict) + and isinstance(self._kwargs.get("parse_dates", []), list) + ) + def _get_csv_read_options(self, csv_engine: str, chunksize: int | None) -> dict[str, Any]: """Build pandas options for Athena CSV or tab-separated results.""" if self.output_location and self.output_location.endswith(".txt"): diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 9d0f1279..7ca3052a 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -19,7 +19,11 @@ from pyathena.model import AthenaQueryExecution from pyathena.pandas.converter import DefaultPandasTypeConverter from pyathena.pandas.cursor import PandasCursor -from pyathena.pandas.result_set import AthenaPandasResultSet, PandasDataFrameIterator +from pyathena.pandas.result_set import ( + AthenaPandasResultSet, + PandasDataFrameIterator, + _read_csv_with_pyarrow, +) from pyathena.util import RetryConfig from tests import ENV from tests.pyathena.conftest import connect @@ -299,6 +303,24 @@ 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 + 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. + with patch( + "pyathena.pandas.result_set._read_csv_with_pyarrow", wraps=_read_csv_with_pyarrow + ) as read_csv_with_pyarrow: + pandas_cursor.execute( + """ + SELECT array_join(repeat('x', 600), '') || chr(10) + || '"' || array_join(repeat('y', 598), '') || '"' AS v + FROM UNNEST(sequence(1, 2000)) AS t(i) + """, + engine="pyarrow", + ) + df = pandas_cursor.as_pandas() + read_csv_with_pyarrow.assert_called_once() + assert df["v"].tolist() == ["x" * 600 + '\n"' + "y" * 598 + '"'] * 2000 + @pytest.mark.parametrize( ("pandas_cursor", "parquet_engine", "chunksize"), [ @@ -1042,6 +1064,43 @@ def test_get_csv_engine_explicit_specification(self): engine = result_set._get_csv_engine() assert engine == "c" + @pytest.mark.parametrize( + ("keep_default_na", "na_values", "kwargs", "expected"), + [ + (False, ("",), {}, True), + (False, ("",), {"dtype": {"a": "str"}, "parse_dates": ["b"]}, True), + (True, ("",), {}, False), + (False, ("", "NA"), {}, False), + (False, np.array([""]), {}, False), + (False, ("",), {"storage_options": None}, False), + (False, ("",), {"on_bad_lines": "skip"}, False), + (False, ("",), {"dtype": "str"}, False), + (False, ("",), {"parse_dates": "d"}, False), + (False, ("",), {"parse_dates": ("d",)}, False), + ], + ids=[ + "default", + "dtype_parse_dates", + "keep_default_na", + "na_values", + "na_values_array", + "storage_options", + "on_bad_lines", + "single_dtype", + "parse_dates_string", + "parse_dates_tuple", + ], + ) + def test_reads_csv_with_pyarrow(self, keep_default_na, na_values, kwargs, expected): + # PyAthena reads the CSV result for the PyArrow engine only with the options + # that _read_csv_with_pyarrow() reproduces; pandas reads it otherwise. + with patch("pyathena.pandas.result_set.AthenaResultSet.__init__"): + result_set = AthenaPandasResultSet.__new__(AthenaPandasResultSet) + result_set._keep_default_na = keep_default_na + result_set._na_values = na_values + result_set._kwargs = kwargs + assert result_set._reads_csv_with_pyarrow() is expected + @pytest.mark.parametrize( ("pandas_cursor", "parquet_engine", "chunksize"), [ diff --git a/tests/pyathena/pandas/test_result_set.py b/tests/pyathena/pandas/test_result_set.py index 84ee9892..ef6022a5 100644 --- a/tests/pyathena/pandas/test_result_set.py +++ b/tests/pyathena/pandas/test_result_set.py @@ -5,16 +5,20 @@ # # SPDX-License-Identifier: MIT +import csv import io -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, PropertyMock, patch import pandas as pd import pytest +from pandas.testing import assert_frame_equal +from pyathena.pandas.converter import DefaultPandasTypeConverter from pyathena.pandas.result_set import ( AthenaPandasResultSet, PandasDataFrameIterator, _no_trunc_date, + _read_csv_with_pyarrow, ) @@ -113,3 +117,164 @@ def test_read_parquet_filesystem(self, execute_kwargs, path, filesystem_kwargs): **execute_kwargs, **filesystem_kwargs, } + + +# A CSV result of Athena with each type the PyArrow engine reads, then a row of NULLs. +_TYPES_CSV = ( + '"ti","si","i","bi","r","d","c","v","ml","arr","m","rw","dt","ts","ts6","tm","iv","nul","u",' + '"empty","na"\n' + '"1","2","3","4","1.5","2.25","ab ","plain","multi\nline ""q"", x","[1, 2]","{k=1}",' + '"{a=1, b=x}","2024-02-29","2024-02-29 23:59:58.123","2024-02-29 23:59:58.123456",' + '"12:34:56.789","2 00:00:00.000",,"589f6631-9c50-4f58-a121-e2608a04fc64","","NA"\n' + ",,,,,,,,,,,,,,,,,,,,\n" +) +_TYPES = { + "ti": "tinyint", + "si": "smallint", + "i": "integer", + "bi": "bigint", + "r": "float", + "d": "double", + "c": "char", + "v": "varchar", + "ml": "varchar", + "arr": "array", + "m": "map", + "rw": "row", + "dt": "date", + "ts": "timestamp", + "ts6": "timestamp", + "tm": "time", + "iv": "interval day to second", + "nul": "unknown", + "u": "uuid", + "empty": "varchar", + "na": "varchar", +} + + +def _pyarrow_read_csv_kwargs(types, tab_separated=False, **kwargs): + """Build the pandas.read_csv() options with AthenaPandasResultSet._get_csv_read_options(). + + Args: + types: The Athena types of the result columns, keyed by column name. + tab_separated: Whether the result is a tab-separated ``.txt`` file. + **kwargs: The pandas.read_csv() options given to ``execute()``. + + Returns: + The options for the PyArrow engine. + """ + with patch("pyathena.pandas.result_set.AthenaResultSet.__init__", return_value=None): + result_set = AthenaPandasResultSet.__new__(AthenaPandasResultSet) + result_set._converter = DefaultPandasTypeConverter() + result_set._keep_default_na = False + result_set._na_values = ("",) + result_set._quoting = 1 + result_set._kwargs = kwargs + description = [(name, type_, None, None, 0, 0, "UNKNOWN") for name, type_ in types.items()] + location = f"s3://bucket/result.{'txt' if tab_separated else 'csv'}" + with ( + patch.object( + AthenaPandasResultSet, + "description", + new_callable=PropertyMock, + return_value=description, + ), + patch.object( + AthenaPandasResultSet, + "output_location", + new_callable=PropertyMock, + return_value=location, + ), + ): + assert result_set._reads_csv_with_pyarrow() + return result_set._get_csv_read_options("pyarrow", None) + + +@pytest.mark.filterwarnings("ignore:Could not infer format") +@pytest.mark.parametrize("infer_string", [True, False]) +@pytest.mark.parametrize( + ("data", "read_csv_kwargs"), + [ + (_TYPES_CSV, _pyarrow_read_csv_kwargs(_TYPES)), + ( + _TYPES_CSV, + _pyarrow_read_csv_kwargs( + _TYPES, + dtype={ + **_pyarrow_read_csv_kwargs(_TYPES)["dtype"], + "ti": "float32", + "v": "category", + "missing": "int64", + }, + ), + ), + (_TYPES_CSV, _pyarrow_read_csv_kwargs(_TYPES, parse_dates=[12, "ts"])), + ( + _TYPES_CSV, + _pyarrow_read_csv_kwargs( + _TYPES, dtype={**_pyarrow_read_csv_kwargs(_TYPES)["dtype"], "dt": "string"} + ), + ), + ( + '"x","d"\n"1","2024-01-01"\n,\n', + _pyarrow_read_csv_kwargs({"x": "integer", "d": "date"}, dtype={"x": None}), + ), + ( + '"x","x","d"\n"1","2","2024-01-01"\n,,\n', + _pyarrow_read_csv_kwargs({"x": "integer", "d": "date"}), + ), + ( + '"v"\n"plain"\n"2024-01-01"\n\n', + _pyarrow_read_csv_kwargs({"v": "varchar"}, parse_dates=["v"]), + ), + ( + "id \tint \t \nname \tstring \t \n", + _pyarrow_read_csv_kwargs({"col_name": "varchar"}, True), + ), + ( + "x\t1\t2024-01-01\n\t\t\ny y\t3\t2024-01-02\n", + _pyarrow_read_csv_kwargs({"a": "varchar", "b": "bigint", "c": "date"}, True), + ), + ], + ids=[ + "types", + "dtype", + "parse_dates", + "dtype_of_date_column", + "dtype_none", + "duplicate_names", + "unparsed_dates", + "tab_separated_extra_fields", + "tab_separated", + ], +) +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"). + 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"])}, + ) + actual = _read_csv_with_pyarrow( + io.BytesIO(data.encode()), + {**read_csv_kwargs, "dtype": dict(read_csv_kwargs["dtype"])}, + ) + assert_frame_equal(actual, expected, check_exact=True) + + +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. + rows = [(i, f'{"x" * 600}\n"{"y" * 598}"' if i % 2 else "z") for i in range(4000)] + buffer = io.StringIO() + writer = csv.writer(buffer, quoting=csv.QUOTE_ALL, lineterminator="\n") + writer.writerow(["id", "v"]) + writer.writerows(rows) + df = _read_csv_with_pyarrow( + io.BytesIO(buffer.getvalue().encode()), + _pyarrow_read_csv_kwargs({"id": "integer", "v": "varchar"}), + ) + assert df["id"].tolist() == [i for i, _ in rows] + assert df["v"].tolist() == [v for _, v in rows]