Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 22 additions & 2 deletions pyiceberg/io/pyarrow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
40 changes: 40 additions & 0 deletions tests/io/test_pyarrow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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`
Expand Down
Loading