diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 55c077261..4390bdbd1 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -343,17 +343,27 @@ def __init__( d[0] for d in description if d[1] in ("time", "time with time zone") ] + import pandas as pd + + # The whole result when it was not read in chunks. + self._df: DataFrame | None = None if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: - df = self._as_pandas() + result = self._as_pandas() trunc_date = _no_trunc_date if self.is_unload else self._trunc_date - self._df_iter = PandasDataFrameIterator(df, trunc_date, self._csv_stream) + if isinstance(result, pd.DataFrame): + self._df = trunc_date(result) + else: + self._df_iter = PandasDataFrameIterator(result, trunc_date, self._csv_stream) elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: - df = self._as_pandas_from_api() - self._df_iter = PandasDataFrameIterator(df, self._trunc_date) + # GetQueryResults values are already converted and need no time truncation. + self._df = self._as_pandas_from_api() else: - import pandas as pd - - self._df_iter = PandasDataFrameIterator(pd.DataFrame(), _no_trunc_date) + self._df = pd.DataFrame() + if self._df is not None: + # A shallow copy keeps assignments to the DataFrame from as_pandas() + # 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) self._iterrows = self._df_iter.iterrows() def _get_parquet_engine(self) -> str: @@ -842,22 +852,30 @@ def as_pandas(self) -> PandasDataFrameIterator | DataFrame: """Return the query results as a DataFrame or an iterator of DataFrame chunks. Returns: - If ``chunksize`` is None, one DataFrame that joins the chunks the result - iterator has not yet yielded (read in chunks when ``auto_optimize_chunksize`` - chose a chunk size), which is the whole result unless rows were already - fetched; otherwise the ``PandasDataFrameIterator`` that yields DataFrame chunks. + If ``chunksize`` is None, the DataFrame of the whole result, the same one + on every call. When ``auto_optimize_chunksize`` chose a chunk size, one + DataFrame that joins the chunks the result iterator has not yet yielded, + which is the whole result only if neither the fetch methods nor + ``iter_chunks()`` read from it before, and a later call returns an empty + DataFrame. If ``chunksize`` is set, the iterator that ``iter_chunks()`` + returns. """ if self._chunksize is None: + if self._df is not None: + return self._df return self._df_iter.as_pandas() - return self._df_iter + return self.iter_chunks() def iter_chunks(self) -> PandasDataFrameIterator: """Iterate over result chunks as pandas DataFrames. This method provides an iterator interface for processing large result sets. - When chunksize is specified, or ``auto_optimize_chunksize`` chose a chunk size - for a large CSV result, it yields DataFrames in chunks for memory-efficient - processing. Otherwise, it yields the entire result as a single DataFrame. + When a CSV result is read in chunks, because chunksize is specified or + ``auto_optimize_chunksize`` chose a chunk size, it yields DataFrames in chunks + for memory-efficient processing. These chunks come from the same iterator as + the fetch methods, so a chunk that one of them reads is not available to the + other. Otherwise, each call returns a new iterator that yields the entire + result as a single DataFrame, and the fetch methods keep their position. Returns: PandasDataFrameIterator that yields pandas DataFrames for each chunk @@ -876,6 +894,8 @@ def iter_chunks(self) -> PandasDataFrameIterator: >>> for df in cursor.iter_chunks(): ... process(df) # Single DataFrame with all data """ + if self._df is not None: + return PandasDataFrameIterator(self._df, _no_trunc_date) return self._df_iter @override @@ -884,6 +904,7 @@ def close(self) -> None: super().close() self._df_iter.close() - self._df_iter = PandasDataFrameIterator(pd.DataFrame(), _no_trunc_date) + self._df = pd.DataFrame() + self._df_iter = PandasDataFrameIterator(self._df, _no_trunc_date) self._iterrows = enumerate([]) self._data_manifest = [] diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index b6a2979bb..35f4614b2 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -254,23 +254,33 @@ def __init__( self._chunksize = chunksize self._kwargs = kwargs - # Build DataFrame iterator (handles both chunked and non-chunked cases) - # Note: _create_dataframe_iterator() calls _as_polars() which may update - # _metadata for unload queries, so we must cache column names AFTER this. + import polars as pl + + # The whole result when it was not read in chunks. + # Note: _as_polars() may update _metadata for unload queries, so the converters + # and column names must be read AFTER it. + self._df: pl.DataFrame | None = None + # Converters for the rows of self._df. GetQueryResults values are already converted. + self._df_converters: dict[str, Callable[[str | None], Any | None]] = {} if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: - self._df_iter = self._create_dataframe_iterator() + if self._chunksize is None: + self._df = self._as_polars() + self._df_converters = self.converters + else: + self._df_iter = self._create_dataframe_iterator() elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: - df = self._as_polars_from_api() - self._df_iter = PolarsDataFrameIterator(df, self.converters, self._get_column_names()) + self._df = self._as_polars_from_api() else: - import polars as pl - + self._df = pl.DataFrame() + if self._df is not None: + # A clone keeps assignments to the DataFrame from as_polars() + # out of the rows that the fetch methods return. self._df_iter = PolarsDataFrameIterator( - pl.DataFrame(), self.converters, self._get_column_names() + self._df.clone(), self._df_converters, self._get_column_names() ) # Cache column names for efficient access in fetchone() - # Must be after _create_dataframe_iterator() which updates _metadata for unload + # Must be after _as_polars() which updates _metadata for unload self._column_names_cache: list[str] = self._get_column_names() self._iterrows = self._df_iter.iterrows() @@ -343,20 +353,12 @@ def _get_column_names(self) -> list[str]: return [d[0] for d in description] def _create_dataframe_iterator(self) -> PolarsDataFrameIterator: - """Create a DataFrame iterator for the result set. + """Create a DataFrame iterator that reads the result file in chunks. Returns: - PolarsDataFrameIterator that handles both chunked and non-chunked cases. + PolarsDataFrameIterator that reads each chunk lazily. """ - if self._chunksize is not None: - # Chunked mode: create lazy iterator - reader: Iterator[pl.DataFrame] | pl.DataFrame = ( - self._iter_parquet_chunks() if self.is_unload else self._iter_csv_chunks() - ) - else: - # Non-chunked mode: load entire DataFrame - reader = self._as_polars() - + reader = self._iter_parquet_chunks() if self.is_unload else self._iter_csv_chunks() return PolarsDataFrameIterator(reader, self.converters, self._get_column_names()) @override @@ -538,12 +540,15 @@ def as_polars(self) -> pl.DataFrame: method for accessing results with PolarsCursor. Note: - When chunksize is set, calling this method will collect all chunks - into a single DataFrame, loading all data into memory. Use - iter_chunks() for memory-efficient processing of large datasets. + When chunksize is set and the result file is read in chunks, calling this + method will collect the chunks that the fetch methods and iter_chunks() + have not yet read into a single DataFrame, loading them all into memory, + and a later call returns an empty DataFrame. Use iter_chunks() for + memory-efficient processing of large datasets. Returns: - Polars DataFrame containing all query results. + Polars DataFrame containing all query results. When the result is not + read in chunks, it is the same DataFrame on every call. Example: >>> cursor = connection.cursor(PolarsCursor) @@ -552,6 +557,8 @@ def as_polars(self) -> pl.DataFrame: >>> print(f"DataFrame has {df.height} rows") >>> filtered = df.filter(pl.col("value") > 100) """ + if self._df is not None: + return self._df return self._df_iter.as_polars() def as_arrow(self) -> Table: @@ -561,7 +568,9 @@ def as_arrow(self) -> Table: interoperability with other Arrow-compatible tools and libraries. Returns: - Apache Arrow Table containing all query results. + Apache Arrow Table containing all query results. When the result file is + read in chunks, it contains the chunks that have not yet been read, as + with as_polars(). Raises: ImportError: If pyarrow is not installed. @@ -573,7 +582,7 @@ def as_arrow(self) -> Table: >>> # Use with other Arrow-compatible libraries """ try: - return self._df_iter.as_polars().to_arrow() + return self.as_polars().to_arrow() except ImportError as e: raise ImportError( "pyarrow is required for as_arrow(). Install it with: pip install pyarrow" @@ -667,8 +676,12 @@ def iter_chunks(self) -> PolarsDataFrameIterator: This method provides an iterator interface for processing large result sets. When chunksize is specified, it yields DataFrames in chunks using lazy - evaluation for memory-efficient processing. When chunksize is not specified, - it yields the entire result as a single DataFrame. + evaluation for memory-efficient processing. These chunks come from the same + iterator as the fetch methods, so a chunk that one of them reads is not + available to the other. When chunksize is not specified, or the result has + no result file to read in chunks, each call returns a new iterator that + yields the entire result as a single DataFrame, and the fetch methods keep + their position. Returns: PolarsDataFrameIterator that yields Polars DataFrames for each chunk @@ -687,6 +700,8 @@ def iter_chunks(self) -> PolarsDataFrameIterator: >>> for df in cursor.iter_chunks(): ... 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 self._df_iter @override @@ -695,5 +710,7 @@ def close(self) -> None: import polars as pl super().close() - self._df_iter = PolarsDataFrameIterator(pl.DataFrame(), {}, []) + self._df_iter.close() + self._df = pl.DataFrame() + self._df_iter = PolarsDataFrameIterator(self._df, {}, []) self._iterrows = iter([]) diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 18faf5d5a..4d558deb2 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -1416,6 +1416,20 @@ def test_pandas_cursor_iter_chunks_without_chunksize(self, pandas_cursor): # Should yield exactly one chunk (the entire DataFrame) assert chunk_count == 1 + def test_pandas_cursor_whole_result_reused(self, pandas_cursor): + """Test that as_pandas() and iter_chunks() do not consume the fetched rows.""" + pandas_cursor.execute("SELECT number FROM (VALUES (1), (2), (3)) AS t(number)") + df = pandas_cursor.as_pandas() + assert df["number"].tolist() == [1, 2, 3] + assert pandas_cursor.as_pandas() is df + df["number"] = 0 + + assert pandas_cursor.fetchone() == (1,) + assert [len(chunk) for chunk in pandas_cursor.iter_chunks()] == [3] + assert [len(chunk) for chunk in pandas_cursor.iter_chunks()] == [3] + assert pandas_cursor.as_pandas() is df + assert pandas_cursor.fetchall() == [(2,), (3,)] + def test_pandas_cursor_chunked_vs_regular_same_data(self, pandas_cursor): """Test that chunked and regular reading produce the same data.""" query = "SELECT * FROM many_rows LIMIT 100" # Use a reasonable size for testing @@ -1515,5 +1529,5 @@ def test_pandas_cursor_iter_chunks_consistency(self, pandas_cursor): indirect=["pandas_cursor"], ) def test_fetch_all_rows(self, pandas_cursor): - pandas_cursor.execute("SELECT 1 AS col") - assert pandas_cursor.fetchall() == [(1,)] + pandas_cursor.execute("SELECT 1 AS col, CAST('12:34:56' AS TIME) AS col_time") + assert pandas_cursor.fetchall() == [(1, datetime(2017, 1, 1, 12, 34, 56).time())] diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index c9da61696..586eba63e 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -474,6 +474,33 @@ def test_iter_chunks_without_chunksize(self, polars_cursor): assert isinstance(chunks[0], pl.DataFrame) assert chunks[0].height == 1 + def test_whole_result_reused(self, polars_cursor): + """Test that as_polars(), as_arrow(), and iter_chunks() do not consume the rows.""" + polars_cursor.execute("SELECT number FROM (VALUES (1), (2), (3)) AS t(number)") + df = polars_cursor.as_polars() + assert df["number"].to_list() == [1, 2, 3] + assert polars_cursor.as_polars() is df + assert polars_cursor.as_arrow().column("number").to_pylist() == [1, 2, 3] + df[0, "number"] = 0 + + assert polars_cursor.fetchone() == (1,) + assert [chunk.height for chunk in polars_cursor.iter_chunks()] == [3] + assert [chunk.height for chunk in polars_cursor.iter_chunks()] == [3] + assert polars_cursor.as_polars() is df + assert polars_cursor.fetchall() == [(2,), (3,)] + + @pytest.mark.parametrize( + "polars_cursor", [{"cursor_kwargs": {"chunksize": 5}}], indirect=["polars_cursor"] + ) + def test_close_stops_chunks(self, polars_cursor): + """Test that closing the result set closes the chunk iterator it returned.""" + polars_cursor.execute("SELECT * FROM many_rows LIMIT 15") + result_set = polars_cursor.result_set + chunks = result_set.iter_chunks() + assert next(chunks).height == 5 + result_set.close() + assert list(chunks) == [] + def test_iter_chunks_many_rows(self): """Test chunked iteration with many rows.""" with contextlib.closing(connect(schema_name=ENV.schema)) as conn: @@ -691,5 +718,5 @@ def test_null_vs_empty_string(self, polars_cursor): indirect=["polars_cursor"], ) def test_fetch_all_rows(self, polars_cursor): - polars_cursor.execute("SELECT 1 AS col") - assert polars_cursor.fetchall() == [(1,)] + polars_cursor.execute("SELECT 1 AS col, CAST('12:34:56' AS TIME) AS col_time") + assert polars_cursor.fetchall() == [(1, datetime(2017, 1, 1, 12, 34, 56).time())]