diff --git a/pyiceberg/io/pyarrow.py b/pyiceberg/io/pyarrow.py index 2dcb8a5795..6bfddfdf18 100644 --- a/pyiceberg/io/pyarrow.py +++ b/pyiceberg/io/pyarrow.py @@ -37,6 +37,7 @@ import re import uuid import warnings +import weakref from abc import ABC, abstractmethod from collections.abc import Callable, Iterable, Iterator from copy import copy @@ -396,11 +397,30 @@ def to_input_file(self) -> PyArrowFile: return self +def _fs_by_scheme_cache(file_io: PyArrowFileIO) -> Callable[[str, str | None], FileSystem]: + """Return a cached FileSystem factory that only weakly references ``file_io``. + + Caching the bound method ``file_io._initialize_fs`` directly would make the cache hold + ``file_io`` while ``file_io`` holds the cache. That reference cycle keeps the FileIO and + its cached filesystems (and their connection pools) alive until the cycle collector runs. + """ + file_io_ref = weakref.ref(file_io) + + @lru_cache + def fs_by_scheme(scheme: str, netloc: str | None = None) -> FileSystem: + io = file_io_ref() + if io is None: + raise ReferenceError("PyArrowFileIO has already been garbage collected") + return io._initialize_fs(scheme, netloc) + + return fs_by_scheme + + class PyArrowFileIO(FileIO): fs_by_scheme: Callable[[str, str | None], FileSystem] def __init__(self, properties: Properties = EMPTY_DICT): - self.fs_by_scheme: Callable[[str, str | None], FileSystem] = lru_cache(self._initialize_fs) + self.fs_by_scheme: Callable[[str, str | None], FileSystem] = _fs_by_scheme_cache(self) super().__init__(properties=properties) @staticmethod @@ -725,7 +745,7 @@ def __getstate__(self) -> dict[str, Any]: def __setstate__(self, state: dict[str, Any]) -> None: """Deserialize the state into a PyArrowFileIO instance.""" self.__dict__ = state - self.fs_by_scheme = lru_cache(self._initialize_fs) + self.fs_by_scheme = _fs_by_scheme_cache(self) def schema_to_pyarrow( diff --git a/tests/io/test_pyarrow.py b/tests/io/test_pyarrow.py index 4d5d4431cb..d8b2d7ccb2 100644 --- a/tests/io/test_pyarrow.py +++ b/tests/io/test_pyarrow.py @@ -15,13 +15,17 @@ # specific language governing permissions and limitations # under the License. # pylint: disable=protected-access,unused-argument,redefined-outer-name +import gc import logging import os +import pickle import sys import tempfile import uuid import warnings +import weakref from collections.abc import Iterator +from contextlib import contextmanager from datetime import date, datetime, timezone from pathlib import Path from typing import Any @@ -3307,6 +3311,42 @@ def test_pyarrow_file_io_fs_by_scheme_cache() -> None: assert pyarrow_file_io.fs_by_scheme.cache_info().hits == 2 # type: ignore +@contextmanager +def _cycle_collector_disabled() -> Iterator[None]: + was_enabled = gc.isenabled() + gc.disable() + try: + yield + finally: + if was_enabled: + gc.enable() + + +def test_pyarrow_file_io_freed_by_refcounting() -> None: + with _cycle_collector_disabled(): + file_io = PyArrowFileIO() + file_io.fs_by_scheme("file", None) + file_io_ref = weakref.ref(file_io) + del file_io + + assert file_io_ref() is None + + +def test_pyarrow_file_io_pickle_round_trip_keeps_cache_and_refcounting() -> None: + file_io = PyArrowFileIO() + file_io.fs_by_scheme("file", None) + + restored = pickle.loads(pickle.dumps(file_io)) + assert isinstance(restored.fs_by_scheme("file", None), LocalFileSystem) + assert restored.fs_by_scheme.cache_info().currsize == 1 + + with _cycle_collector_disabled(): + restored_ref = weakref.ref(restored) + del restored + + assert restored_ref() is None + + def test_pyarrow_io_new_input_multi_region(caplog: Any) -> None: # It's better to set up multi-region minio servers for an integration test once `endpoint_url` argument # becomes available for `resolve_s3_region`