diff --git a/pyathena/filesystem/s3_async.py b/pyathena/filesystem/s3_async.py index eb1d78b92..4157c26a6 100644 --- a/pyathena/filesystem/s3_async.py +++ b/pyathena/filesystem/s3_async.py @@ -123,6 +123,8 @@ def __init__( allow_bucket_creation=allow_bucket_creation, allow_bucket_deletion=allow_bucket_deletion, version_aware=version_aware, + # fsspec caches the AioS3FileSystem itself when caching is wanted. + skip_instance_cache=True, **kwargs, ) # Share dircache for cache coherence between async and sync instances diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 55c077261..10576a678 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -8,7 +8,7 @@ from collections.abc import Callable, Iterable, Iterator from contextlib import ExitStack from functools import partial -from io import BufferedReader, StringIO, TextIOWrapper +from io import BufferedReader, IOBase, StringIO, TextIOWrapper from multiprocessing import cpu_count from typing import ( TYPE_CHECKING, @@ -72,7 +72,7 @@ def __init__( self, reader: TextFileReader | DataFrame, trunc_date: Callable[[DataFrame], DataFrame], - csv_stream: TextIOWrapper | None = None, + csv_stream: IOBase | None = None, ) -> None: """Initialize the iterator. @@ -335,7 +335,7 @@ def __init__( self._data_manifest: list[str] = [] self._kwargs = kwargs self._fs = self._create_s3_file_system() - self._csv_stream: TextIOWrapper | None = None + self._csv_stream: IOBase | None = None # Cache time column names for efficient _trunc_date processing description = self.description if self.description else [] @@ -465,11 +465,14 @@ def _create_s3_file_system(self): """ from pyathena.filesystem.s3 import S3FileSystem + # Not cached by fsspec so that the connection and the dircache are + # released with the result set. return S3FileSystem( connection=self.connection, default_block_size=self._block_size, default_cache_type=self._cache_type, max_workers=self._max_workers, + skip_instance_cache=True, ) @property @@ -558,14 +561,22 @@ def _read_csv(self) -> TextFileReader | DataFrame: try: with ExitStack() as stack: - source: str | TextIOWrapper = self.output_location + source: str | IOBase = self.output_location binary_columns = self._configure_binary_csv_read(read_csv_kwargs, pd.read_csv) if binary_columns: - storage_options = read_csv_kwargs.pop("storage_options", None) or {} - self._csv_stream = stack.enter_context( + # Given storage_options, even None, open the file through fsspec + # as pandas does. + storage_options = None + if "storage_options" in read_csv_kwargs: + storage_options = read_csv_kwargs.pop("storage_options") or {} + source = self._csv_stream = stack.enter_context( self._open_binary_csv_stream(binary_columns, storage_options) ) - source = self._csv_stream + elif "storage_options" not in read_csv_kwargs: + # With storage_options, pandas opens the file through fsspec. + source = self._csv_stream = stack.enter_context( + self._fs.open(self.output_location, mode="rb") + ) result = pd.read_csv(source, **read_csv_kwargs) if not isinstance(result, pd.DataFrame): # The chunk iterator takes ownership of the stream. @@ -608,12 +619,6 @@ def _get_csv_read_options(self, csv_engine: str, chunksize: int | None) -> dict[ "keep_default_na": self._keep_default_na, "na_values": self._na_values, "quoting": self._quoting, - "storage_options": { - "connection": self.connection, - "default_block_size": self._block_size, - "default_cache_type": self._cache_type, - "max_workers": self._max_workers, - }, "chunksize": chunksize, "engine": csv_engine, } @@ -736,19 +741,17 @@ def _configure_binary_csv_read( return binary_columns def _open_binary_csv_stream( - self, binary_columns: set[int], storage_options: dict[str, Any] + self, binary_columns: set[int], storage_options: dict[str, Any] | None ) -> TextIOWrapper: """Open a stream that preserves binary NULL fields and original CSV newlines.""" + text_options: dict[str, Any] = {"mode": "rt", "encoding": "utf-8", "newline": ""} with ExitStack() as stack: - source = stack.enter_context( - filesystem_open( - self.output_location, - mode="rt", - encoding="utf-8", - newline="", - **storage_options, + if storage_options is None: + source = stack.enter_context(self._fs.open(self.output_location, **text_options)) + else: + source = stack.enter_context( + filesystem_open(self.output_location, **text_options, **storage_options) ) - ) reader = stack.enter_context(BinaryCSVReader(source, binary_columns)) buffer = stack.enter_context(BufferedReader(reader)) stream = TextIOWrapper(buffer, encoding="utf-8", newline="") @@ -765,7 +768,9 @@ def _read_parquet(self, engine) -> DataFrame: self._unload_location = "/".join(self._data_manifest[0].split("/")[:-1]) + "/" if engine == "pyarrow": - unload_location = self._unload_location + # pyarrow takes the path without the scheme with an fsspec filesystem. + bucket, key = parse_output_location(self._unload_location) + unload_location = f"{bucket}/{key}" kwargs = { "use_threads": True, } @@ -777,12 +782,7 @@ def _read_parquet(self, engine) -> DataFrame: return pd.read_parquet( unload_location, engine=self._engine, - storage_options={ - "connection": self.connection, - "default_block_size": self._block_size, - "default_cache_type": self._cache_type, - "max_workers": self._max_workers, - }, + filesystem=self._fs, **kwargs, ) except Exception as e: diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index b6a2979bb..1fced8450 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -289,6 +289,9 @@ def _csv_storage_options(self) -> dict[str, Any]: "default_block_size": self._block_size, "default_cache_type": self._cache_type, "max_workers": self._max_workers, + # Not cached by fsspec so that the connection and the dircache are + # released with the result set. + "skip_instance_cache": True, } @property diff --git a/pyathena/s3fs/result_set.py b/pyathena/s3fs/result_set.py index 151bf82fc..74dff6c7b 100644 --- a/pyathena/s3fs/result_set.py +++ b/pyathena/s3fs/result_set.py @@ -131,9 +131,12 @@ def __init__( def _create_s3_file_system(self) -> AbstractFileSystem: """Create S3FileSystem using connection settings.""" + # Not cached by fsspec so that the connection and the dircache are + # released with the result set. return self._filesystem_class( connection=self.connection, default_block_size=self._block_size, + skip_instance_cache=True, ) def _init_csv_reader(self) -> None: diff --git a/tests/pyathena/filesystem/test_s3_async.py b/tests/pyathena/filesystem/test_s3_async.py index 59604bb97..cf9441b5c 100644 --- a/tests/pyathena/filesystem/test_s3_async.py +++ b/tests/pyathena/filesystem/test_s3_async.py @@ -503,6 +503,20 @@ def upload_part_copy(**kw): sync_fs._complete_multipart_upload.call_args.kwargs["RequestPayer"] == "requester" ) + def test_internal_file_system_not_cached(self): + # GH-978: the internal S3FileSystem was kept in the fsspec instance + # cache, so skip_instance_cache=True instances shared it. + connection = mock.MagicMock() + fs1 = AioS3FileSystem(connection=connection, skip_instance_cache=True) + fs2 = AioS3FileSystem(connection=connection, skip_instance_cache=True) + assert fs1._sync_fs is not fs2._sync_fs + assert fs1.dircache is not fs2.dircache + + fs3 = AioS3FileSystem(connection=connection) + assert AioS3FileSystem(connection=connection) is fs3 + for fs in (fs1, fs2, fs3): + assert fs._sync_fs not in S3FileSystem._cache.values() + @pytest.fixture(scope="class") def fs(self, request): if not hasattr(request, "param"): diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 18faf5d5a..1ebbe2d25 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -15,11 +15,13 @@ from pandas.io.parsers import TextFileReader from pyathena.error import DatabaseError, ProgrammingError +from pyathena.filesystem.s3 import S3FileSystem from pyathena.pandas.converter import DefaultPandasTypeConverter from pyathena.pandas.cursor import PandasCursor from pyathena.pandas.result_set import AthenaPandasResultSet, PandasDataFrameIterator from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.util import cached_file_systems class TestPandasCursor: @@ -230,6 +232,65 @@ def test_fetchone(self, pandas_cursor, parquet_engine, chunksize): assert pandas_cursor.rownumber == 1 assert pandas_cursor.fetchone() is None + @pytest.mark.parametrize( + ("pandas_cursor", "chunksize"), + [ + ({"cursor_kwargs": {"unload": False}}, None), + ({"cursor_kwargs": {"unload": False}}, 1_000), + ({"cursor_kwargs": {"unload": True}}, None), + ], + indirect=["pandas_cursor"], + ) + def test_result_set_file_system(self, pandas_cursor, chunksize): + # GH-978: the filesystems that read the results were kept in the fsspec + # instance cache with the connection, so the connection was never freed. + # The result set reads through its own filesystem instead of creating + # another one from storage_options. + with patch.object( + S3FileSystem, "__init__", autospec=True, side_effect=S3FileSystem.__init__ + ) as init: + pandas_cursor.execute("SELECT * FROM one_row", chunksize=chunksize) + assert pandas_cursor.fetchall() == [(1,)] + assert init.call_count == 1 + assert not cached_file_systems(pandas_cursor.connection) + if not pandas_cursor.result_set.is_unload: + assert pandas_cursor.result_set._csv_stream.closed + + @pytest.mark.parametrize( + ("query", "expected", "binary"), + [ + ("SELECT * FROM one_row", [(1,)], False), + ("SELECT X'01' AS value", [(b"\x01",)], True), + ], + ids=["plain", "binary"], + ) + @pytest.mark.parametrize("with_options", [False, True], ids=["none", "options"]) + def test_csv_storage_options(self, pandas_cursor, query, expected, binary, with_options): + # Given storage_options, even None, the CSV output is opened through fsspec + # with them, as pandas does, instead of the result set's filesystem. + storage_options = ( + { + "connection": pandas_cursor.connection, + "default_cache_type": "none", + "skip_instance_cache": True, + } + if with_options + else None + ) + with patch.object( + S3FileSystem, "open", autospec=True, side_effect=S3FileSystem.open + ) as open_: + pandas_cursor.execute(query, storage_options=storage_options) + assert pandas_cursor.fetchall() == expected + file_systems = [c.args[0] for c in open_.call_args_list] + assert not [fs for fs in file_systems if fs is pandas_cursor.result_set._fs] + if with_options: + assert file_systems + assert all(fs.default_cache_type == "none" for fs in file_systems) + if not binary: + # pandas opens and closes the file itself. + assert pandas_cursor.result_set._csv_stream is None + @pytest.mark.parametrize( ("pandas_cursor", "parquet_engine", "chunksize"), [ diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index c9da61696..343cefb92 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -21,6 +21,7 @@ from pyathena.polars.result_set import AthenaPolarsResultSet from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.util import cached_file_systems class TestPolarsCursor: @@ -36,6 +37,13 @@ def test_fetchone(self, polars_cursor): assert polars_cursor.rownumber == 1 assert polars_cursor.fetchone() is None + def test_result_set_file_system_not_cached(self, polars_cursor): + # GH-978: the filesystem that read the CSV results was kept in the fsspec + # instance cache with the connection, so the connection was never freed. + polars_cursor.execute("SELECT * FROM one_row") + assert polars_cursor.fetchall() == [(1,)] + assert not cached_file_systems(polars_cursor.connection) + @pytest.mark.parametrize( "polars_cursor", [{"cursor_kwargs": {"unload": False}}, {"cursor_kwargs": {"unload": True}}], diff --git a/tests/pyathena/s3fs/test_cursor.py b/tests/pyathena/s3fs/test_cursor.py index 40edaa072..777303f4e 100644 --- a/tests/pyathena/s3fs/test_cursor.py +++ b/tests/pyathena/s3fs/test_cursor.py @@ -14,6 +14,7 @@ from pyathena.s3fs.result_set import AthenaS3FSResultSet from tests import ENV from tests.pyathena.conftest import connect +from tests.pyathena.util import cached_file_systems class TestS3FSCursor: @@ -24,6 +25,13 @@ def test_fetchone(self, s3fs_cursor): assert s3fs_cursor.rownumber == 1 assert s3fs_cursor.fetchone() is None + def test_result_set_file_system_not_cached(self, s3fs_cursor): + # GH-978: the filesystem that read the results was kept in the fsspec + # instance cache with the connection, so the connection was never freed. + s3fs_cursor.execute("SELECT * FROM one_row") + assert s3fs_cursor.fetchall() == [(1,)] + assert not cached_file_systems(s3fs_cursor.connection) + def test_fetchmany(self, s3fs_cursor): s3fs_cursor.execute("SELECT * FROM many_rows LIMIT 15") assert len(s3fs_cursor.fetchmany(10)) == 10 diff --git a/tests/pyathena/util.py b/tests/pyathena/util.py index 4c6a9bf4d..b2e0cf400 100644 --- a/tests/pyathena/util.py +++ b/tests/pyathena/util.py @@ -15,6 +15,8 @@ from jinja2 import Environment, FileSystemLoader from sqlalchemy import types +from pyathena.filesystem.s3 import S3FileSystem +from pyathena.filesystem.s3_async import AioS3FileSystem from pyathena.glue import GlueMetadataClient from pyathena.model import AthenaCalculationExecutionStatus, AthenaQueryExecution @@ -28,6 +30,16 @@ def read_query(name, **kwargs): return [q.strip() for q in template.render(**kwargs).split(";") if q and q.strip()] +def cached_file_systems(connection): + """Return the filesystems in the fsspec instance cache that hold the connection.""" + return [ + fs + for cls in (S3FileSystem, AioS3FileSystem) + for fs in cls._cache.values() + if fs.storage_options.get("connection") is connection + ] + + METADATA_OPERATIONS = ("get_table_metadata", "list_table_metadata", "list_databases")