diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 54a864e0..d700fa72 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -27,7 +27,7 @@ from pyathena.util import RetryConfig, override, parse_output_location if TYPE_CHECKING: - from pandas import DataFrame + from pandas import DataFrame, Index, Series from pandas.io.parsers import TextFileReader from pyathena.connection import Connection @@ -43,6 +43,75 @@ def _no_trunc_date(df: DataFrame) -> DataFrame: return df +class _JSONConverter: + """A json converter for ``pandas.read_csv()`` that keeps NULL from making values floats. + + pandas infers a dtype from the values that a converter returns, so JSON numbers + with NULL become float64. This converter returns ``NULL`` in place of None, which + keeps the column object, and ``restore()`` puts None back after reading. + """ + + NULL: ClassVar[object] = object() + + __slots__ = ("_converter", "_has_null") + + def __init__(self, converter: Callable[[str | None], Any]) -> None: + """Wrap a json conversion function. + + Args: + converter: The conversion function, which returns None for NULL. + """ + self._converter = converter + self._has_null = False + + def __call__(self, value: str | None) -> Any: + """Convert a CSV value. + + Args: + value: The value as text. + + Returns: + The converted value, or ``NULL`` in place of None. + """ + converted = self._converter(value) + if converted is None: + self._has_null = True + return self.NULL + return converted + + def restore(self, df: DataFrame, name: Any) -> None: + """Put None back in place of ``NULL`` in a DataFrame that ``read_csv()`` returned. + + Only does anything if this converter returned ``NULL`` since the last call, + so call it for each DataFrame or chunk right after reading it. + + Args: + df: The DataFrame or chunk. + name: The name of the column that this converter converted, which can + also be an index level. + """ + if not self._has_null: + return + self._has_null = False + + import pandas as pd + + index = df.index + if name in df.columns: + df[name] = pd.Series(self._restore_values(df[name]), index=index, dtype=object) + elif isinstance(index, pd.MultiIndex) and name in index.names: + levels = [index.get_level_values(i) for i in range(index.nlevels)] + i = index.names.index(name) + levels[i] = pd.Index(self._restore_values(levels[i]), dtype=object, name=name) + df.index = pd.MultiIndex.from_arrays(levels, names=index.names) + elif index.name == name: + df.index = pd.Index(self._restore_values(index), dtype=object, name=name) + + @classmethod + def _restore_values(cls, values: Series | Index) -> list[Any]: + return [None if v is cls.NULL else v for v in values.to_numpy()] + + class PandasDataFrameIterator(abc.Iterator): # type: ignore[type-arg] """Iterator for chunked DataFrame results from Athena queries. @@ -78,7 +147,7 @@ def __init__( Args: reader: Either a TextFileReader (for chunked) or a single DataFrame. - trunc_date: Function to apply date truncation to each chunk. + trunc_date: Function to apply to each chunk, such as date truncation. csv_stream: Optional CSV stream owned and closed by this iterator. """ from pandas import DataFrame @@ -254,6 +323,7 @@ class AthenaPandasResultSet(AthenaResultSet): AUTO_CHUNK_SIZE_LARGE: int = 100_000 AUTO_CHUNK_SIZE_MEDIUM: int = 50_000 + _INTEGER_TYPES: ClassVar[tuple[str, ...]] = ("tinyint", "smallint", "integer", "bigint") _PARSE_DATES: ClassVar[list[str]] = [ "date", "time", @@ -338,6 +408,8 @@ def __init__( self._kwargs = kwargs self._fs = self._create_s3_file_system() self._csv_stream: IOBase | None = None + # 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 description = self.description if self.description else [] @@ -349,7 +421,7 @@ def __init__( self._df: DataFrame | None = None if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: result = self._as_pandas() - trunc_date = _no_trunc_date if self.is_unload else self._trunc_date + trunc_date = _no_trunc_date if self.is_unload else self._finish_csv_frame if isinstance(result, pd.DataFrame): self._df = trunc_date(result) else: @@ -516,6 +588,39 @@ def parse_dates(self) -> list[Any | None]: description = self.description if self.description else [] return [d[0] for d in description if d[1] in self._PARSE_DATES] + def _finish_csv_frame(self, df: DataFrame) -> DataFrame: + """Finish a DataFrame read from the CSV result file. + + Puts None back in the json columns and truncates the time columns. + + Args: + df: The DataFrame or chunk that ``pandas.read_csv()`` returned. + + Returns: + The same DataFrame. + """ + for name, converter in self._csv_converters.items(): + if isinstance(converter, _JSONConverter): + converter.restore(df, name) + return self._trunc_date(df) + + def _get_csv_converter(self, type_: str) -> Callable[[str | None], Any]: + """Get the converter that ``pandas.read_csv()`` applies to a column type. + + json columns get a ``_JSONConverter``, so that NULL does not make pandas infer + a numeric dtype, and ``_finish_csv_frame()`` puts None back. + + Args: + type_: The Athena type of the column. + + Returns: + The conversion function. + """ + converter = self._converter.get(type_) + if type_ == "json": + return _JSONConverter(converter) + return converter + def _trunc_date(self, df: DataFrame) -> DataFrame: if self._time_columns: # A NULL is None, as with the GetQueryResults fallback and the other types. @@ -576,6 +681,7 @@ 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, pd.read_csv) + self._csv_converters = read_csv_kwargs.get("converters") or {} if binary_columns: # Given storage_options, even None, open the file through fsspec # as pandas does. @@ -626,7 +732,11 @@ def _get_csv_read_options(self, csv_engine: str, chunksize: int | None) -> dict[ "header": header, "names": names, "dtype": self.dtypes, - "converters": self.converters, + "converters": { + d[0]: self._get_csv_converter(d[1]) + for d in self.description or [] + if d[1] in self._converter.mappings + }, "parse_dates": self.parse_dates, "skip_blank_lines": False, "keep_default_na": self._keep_default_na, @@ -740,7 +850,7 @@ def _configure_binary_csv_read( if len(column_names) != len(description): return set() converters = { - name: self._converter.get(d[1]) + 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 } @@ -848,7 +958,23 @@ def _as_pandas_from_api(self, converter: Converter | None = None) -> DataFrame: return pd.DataFrame() description = self.description if self.description else [] columns = [d[0] for d in description] - return pd.DataFrame(self._rows_to_columnar(rows, columns)) + columnar = self._rows_to_columnar(rows, columns) + # 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: + 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() + } + ) def as_pandas(self) -> PandasDataFrameIterator | DataFrame: """Return the query results as a DataFrame or an iterator of DataFrame chunks. diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index d4d78c9c..9d0f1279 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -26,6 +26,12 @@ from tests.pyathena.util import CONVERTED_VALUES_QUERY, CONVERTED_VALUES_ROW, cached_file_systems +def _pandas_converter_without_bigint_dtype(): + converter = DefaultPandasTypeConverter() + del converter.types["bigint"] + return converter + + class TestPandasCursor: @pytest.mark.parametrize( ("engine", "chunksize"), [("auto", None), ("c", 2), ("python", 2), ("pyarrow", None)] @@ -1671,3 +1677,106 @@ def test_read_options(self, execute_kwargs): kwargs = result_set_class.call_args.kwargs expected = {**cursor_kwargs, **execute_kwargs} assert {key: kwargs[key] for key in expected} == expected + + @pytest.mark.parametrize( + "pandas_cursor", + [ + pytest.param({"converter": _pandas_converter_without_bigint_dtype()}, id="default"), + pytest.param( + { + "work_group": ENV.managed_work_group, + "s3_staging_dir": "", + "converter": _pandas_converter_without_bigint_dtype(), + }, + id="managed", + marks=pytest.mark.skipif( + not ENV.managed_work_group, + reason="AWS_ATHENA_MANAGED_WORKGROUP not set", + ), + ), + ], + indirect=["pandas_cursor"], + ) + def test_integer_without_dtype(self, pandas_cursor): + pandas_cursor.execute( + "SELECT * FROM (VALUES BIGINT '1', NULL) AS t(col_bigint) ORDER BY col_bigint" + ) + df = pandas_cursor.as_pandas() + assert df["col_bigint"].dtype == np.float64 + assert df["col_bigint"].iloc[0] == 1.0 + assert math.isnan(df["col_bigint"].iloc[1]) + + @pytest.mark.parametrize( + ("pandas_cursor", "kwargs"), + [ + pytest.param({}, {}, id="default"), + pytest.param({}, {"chunksize": 1}, id="chunked"), + pytest.param({}, {"index_col": "col_json"}, id="index"), + pytest.param({}, {"index_col": ["col_int", "col_json"]}, id="multi_index"), + pytest.param( + {}, + { + "usecols": [ + "col_int", + "col_bigint", + "col_json", + "col_binary", + "col_json_not_null", + "col_time", + ] + }, + id="usecols", + ), + 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_integer_and_json_with_null(self, pandas_cursor, kwargs): + pandas_cursor.execute( + """ + SELECT * FROM (VALUES + (1, BIGINT '9007199254740993', json_parse('9007199254740993'), X'01', + json_parse('9007199254740995'), CAST('01:02:03' AS TIME)), + (2, NULL, NULL, NULL, json_parse('9007199254740997'), CAST('04:05:06' AS TIME)) + ) AS t(col_int, col_bigint, col_json, col_binary, col_json_not_null, col_time) + ORDER BY col_int + """, + **kwargs, + ) + df = pandas_cursor.as_pandas() + if "chunksize" in kwargs: + df = pd.concat([df.get_chunk(1), *df]) + + def column(name): + if name in df.index.names: + return df.index.get_level_values(name) + return df[name] + + assert column("col_int").dtype == pd.Int64Dtype() + assert column("col_bigint").dtype == pd.Int64Dtype() + assert column("col_json").dtype == np.object_ + # A json column without NULL keeps the dtype that pandas infers. + assert column("col_json_not_null").dtype == np.int64 + assert column("col_bigint").tolist() == [9007199254740993, pd.NA] + if isinstance(df.index, pd.MultiIndex): + # A MultiIndex level holds None as NaN. + json_values = column("col_json").tolist() + assert json_values[0] == 9007199254740993 + assert isinstance(json_values[0], int) + assert pd.isna(json_values[1]) + else: + assert column("col_json").tolist() == [9007199254740993, None] + assert column("col_json_not_null").tolist() == [9007199254740995, 9007199254740997] + assert column("col_binary").tolist() == [b"\x01", None] + assert column("col_time").tolist() == [ + datetime(2017, 1, 1, 1, 2, 3).time(), + datetime(2017, 1, 1, 4, 5, 6).time(), + ]