This is an automated email from the ASF dual-hosted git repository.

JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git


The following commit(s) were added to refs/heads/master by this push:
     new b6bc5910bc [python] Cache file-format metadata across reads (#8731)
b6bc5910bc is described below

commit b6bc5910bcc92c1b079f12b31af3a8f0039285a5
Author: XiaoHongbo <[email protected]>
AuthorDate: Fri Jul 31 10:18:38 2026 +0800

    [python] Cache file-format metadata across reads (#8731)
---
 docs/docs/pypaimon/pytorch.md                      |  20 +
 paimon-python/pypaimon/common/options/config.py    |  11 +
 .../pypaimon/read/reader/format_pyarrow_reader.py  | 201 ++++++++-
 .../pypaimon/tests/parquet_metadata_cache_test.py  | 479 +++++++++++++++++++++
 4 files changed, 708 insertions(+), 3 deletions(-)

diff --git a/docs/docs/pypaimon/pytorch.md b/docs/docs/pypaimon/pytorch.md
index 0f7f7bbdef..6ab0af5173 100644
--- a/docs/docs/pypaimon/pytorch.md
+++ b/docs/docs/pypaimon/pytorch.md
@@ -59,6 +59,26 @@ when it is false, it will read the full amount of data into 
memory.
 
 **`prefetch_concurrency`** (default: 1): When streaming is true, number of 
threads used for parallel prefetch within each DataLoader worker. Set to a 
value greater than 1 to partition splits across threads and increase read 
throughput. Has no effect when streaming is false.
 
+## File Format Metadata Cache
+
+Reusable PyArrow Dataset metadata is cached across reads. Configure its 
estimated
+size limit in the catalog options:
+
+```python
+catalog = CatalogFactory.create({
+    "warehouse": "file:///path/to/warehouse",
+    "file-format.metadata-cache.max-size": "50 mb",
+})
+table = catalog.get_table("database.table")
+read_builder = table.new_read_builder()
+```
+
+The default limit is 50 MB; set it to `0 b` to disable and clear the cache. The
+cache is local to each process and benefits workers reused with
+`DataLoader(..., persistent_workers=True)`. The cache uses a conservative
+per-entry memory estimate and an internal entry-count safeguard; actual native
+PyArrow memory may still be higher. The cache assumes immutable Paimon data 
files.
+
 ## Shuffle
 
 PyPaimon supports streaming shuffle for PyTorch `IterableDataset`. The shuffle
diff --git a/paimon-python/pypaimon/common/options/config.py 
b/paimon-python/pypaimon/common/options/config.py
index 604d62911a..2573782c42 100644
--- a/paimon-python/pypaimon/common/options/config.py
+++ b/paimon-python/pypaimon/common/options/config.py
@@ -15,6 +15,7 @@
 # specific language governing permissions and limitations
 # under the License.
 
+from pypaimon.common.memory_size import MemorySize
 from pypaimon.common.options.config_options import ConfigOptions
 
 
@@ -78,6 +79,16 @@ class CatalogOptions:
     METASTORE = 
ConfigOptions.key("metastore").string_type().default_value("filesystem").with_description(
         "Metastore type")
     WAREHOUSE = 
ConfigOptions.key("warehouse").string_type().no_default_value().with_description("Warehouse
 path")
+    FILE_FORMAT_METADATA_CACHE_MAX_SIZE = (
+        ConfigOptions.key("file-format.metadata-cache.max-size")
+        .memory_type()
+        .default_value(MemorySize.of_mebi_bytes(50))
+        .with_description(
+            "Maximum estimated size of reusable PyArrow Dataset metadata "
+            "cached in the current process. Set to 0 to disable and clear "
+            "the cache."
+        )
+    )
     TOKEN_PROVIDER = 
ConfigOptions.key("token.provider").string_type().no_default_value().with_description(
         "Token provider")
     TOKEN = 
ConfigOptions.key("token").string_type().no_default_value().with_description("Authentication
 token")
diff --git a/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py 
b/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py
index 565ec1f5b1..c84a1383ad 100644
--- a/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py
+++ b/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py
@@ -15,14 +15,19 @@
 # specific language governing permissions and limitations
 # under the License.
 
-from collections import deque
-from typing import Any, Deque, Dict, Iterator, List, Optional, Set, Tuple
+import os
+import sys
+import threading
+from collections import OrderedDict, deque
+from concurrent.futures import Future
+from typing import Any, Callable, Deque, Dict, Iterator, List, Optional, Set, 
Tuple
 
 import pyarrow as pa
 import pyarrow.dataset as ds
 from pyarrow import RecordBatch
 
 from pypaimon.common.file_io import FileIO
+from pypaimon.common.options.config import CatalogOptions
 from pypaimon.common.options.core_options import CoreOptions
 from pypaimon.data.variant_shredding import (
     VariantSchema,
@@ -43,6 +48,194 @@ from pypaimon.schema.data_types import (
 from pypaimon.table.special_fields import SpecialFields
 
 
+_DEFAULT_FILE_FORMAT_METADATA_CACHE_MAX_SIZE = 50 * 1024 * 1024
+_FILE_FORMAT_METADATA_CACHE_MAX_ENTRIES = 4096
+_FILE_FORMAT_METADATA_CACHE_MIN_ENTRY_SIZE = 8 * 1024
+_FILE_FORMAT_METADATA_CACHE_CONTAINER_OVERHEAD = 256
+
+
+class _FilesystemIdentity:
+    def __init__(self, filesystem):
+        self.filesystem = filesystem
+
+    def __hash__(self):
+        return id(self.filesystem)
+
+    def __eq__(self, other):
+        return (
+            isinstance(other, _FilesystemIdentity)
+            and self.filesystem is other.filesystem
+        )
+
+
+class _FileFormatDatasetCache:
+    def __init__(
+            self,
+            max_size: int,
+            max_entries: int = _FILE_FORMAT_METADATA_CACHE_MAX_ENTRIES):
+        self.max_size = max_size
+        self.max_entries = max_entries
+        self.estimated_size = 0
+        self._entries = OrderedDict()
+        self._loads = {}
+        self._lock = threading.Lock()
+
+    def get_or_load(self, key: Tuple[Any, str, str], loader: Callable[[], Any],
+                    size_estimator: Callable[[Any], Optional[int]]):
+        with self._lock:
+            entry = self._entries.get(key)
+            if entry is not None:
+                self._entries.move_to_end(key)
+                return entry[0]
+
+            future = self._loads.get(key)
+            if future is None:
+                future = Future()
+                self._loads[key] = future
+                should_load = True
+            else:
+                should_load = False
+
+        if not should_load:
+            return future.result()
+
+        try:
+            dataset = loader()
+            estimated_size = size_estimator(dataset)
+        except BaseException as exception:
+            future.set_exception(exception)
+            with self._lock:
+                self._loads.pop(key, None)
+            raise
+
+        with self._lock:
+            if estimated_size is not None:
+                estimated_size = max(1, estimated_size)
+                self._entries[key] = (dataset, estimated_size)
+                self.estimated_size += estimated_size
+                self._entries.move_to_end(key)
+                self._evict()
+        future.set_result(dataset)
+        with self._lock:
+            self._loads.pop(key, None)
+        return dataset
+
+    def resize(self, max_size: int):
+        with self._lock:
+            self.max_size = max_size
+            self._evict()
+
+    def _evict(self):
+        while (
+                self.estimated_size > self.max_size
+                or len(self._entries) > self.max_entries):
+            _, (_, evicted_size) = self._entries.popitem(last=False)
+            self.estimated_size -= evicted_size
+
+
+_FILE_FORMAT_DATASET_CACHE = None
+_FILE_FORMAT_DATASET_CACHE_LOCK = threading.Lock()
+_FILE_FORMAT_DATASET_CACHE_PID = os.getpid()
+
+
+def _ensure_file_format_dataset_cache_process():
+    global _FILE_FORMAT_DATASET_CACHE
+    global _FILE_FORMAT_DATASET_CACHE_LOCK
+    global _FILE_FORMAT_DATASET_CACHE_PID
+    current_pid = os.getpid()
+    if current_pid != _FILE_FORMAT_DATASET_CACHE_PID:
+        _FILE_FORMAT_DATASET_CACHE = None
+        _FILE_FORMAT_DATASET_CACHE_LOCK = threading.Lock()
+        _FILE_FORMAT_DATASET_CACHE_PID = current_pid
+
+
+def _file_format_dataset_cache(max_size: int) -> _FileFormatDatasetCache:
+    global _FILE_FORMAT_DATASET_CACHE
+    _ensure_file_format_dataset_cache_process()
+    with _FILE_FORMAT_DATASET_CACHE_LOCK:
+        if _FILE_FORMAT_DATASET_CACHE is None:
+            _FILE_FORMAT_DATASET_CACHE = _FileFormatDatasetCache(max_size)
+        else:
+            _FILE_FORMAT_DATASET_CACHE.resize(max_size)
+        return _FILE_FORMAT_DATASET_CACHE
+
+
+def _reset_file_format_dataset_cache():
+    global _FILE_FORMAT_DATASET_CACHE
+    _ensure_file_format_dataset_cache_process()
+    with _FILE_FORMAT_DATASET_CACHE_LOCK:
+        _FILE_FORMAT_DATASET_CACHE = None
+
+
+def _estimate_file_format_dataset_size(dataset, file_format: str) -> 
Optional[int]:
+    try:
+        if file_format == 'parquet':
+            footer_size = 0
+            for fragment in dataset.get_fragments():
+                metadata = fragment.metadata
+                if metadata is not None:
+                    footer_size += int(metadata.serialized_size)
+            if footer_size > 0:
+                return footer_size
+        return int(dataset.schema.serialize().size)
+    except Exception:
+        return None
+
+
+def _estimate_file_format_cache_entry_size(
+        key: Tuple[Any, str, str],
+        dataset,
+        file_format: str) -> Optional[int]:
+    metadata_size = _estimate_file_format_dataset_size(dataset, file_format)
+    if metadata_size is None:
+        return None
+
+    # PyArrow does not expose the native size retained by Dataset and Fragment
+    # objects. Account for all visible Python objects and apply a conservative
+    # floor so tiny files cannot turn a byte-bounded cache into an effectively
+    # unbounded object cache.
+    visible_size = (
+        metadata_size
+        + sys.getsizeof(key)
+        + sys.getsizeof(key[0])
+        + sys.getsizeof(key[1])
+        + sys.getsizeof(key[2])
+        + sys.getsizeof(dataset)
+        + sys.getsizeof((dataset, metadata_size))
+        + _FILE_FORMAT_METADATA_CACHE_CONTAINER_OVERHEAD
+    )
+    return max(_FILE_FORMAT_METADATA_CACHE_MIN_ENTRY_SIZE, visible_size)
+
+
+def _file_format_metadata_cache_max_size(file_io: FileIO) -> int:
+    properties = getattr(file_io, 'properties', None)
+    if properties is None:
+        return _DEFAULT_FILE_FORMAT_METADATA_CACHE_MAX_SIZE
+    return properties.get(
+        CatalogOptions.FILE_FORMAT_METADATA_CACHE_MAX_SIZE).get_bytes()
+
+
+def _file_format_dataset(file_io: FileIO, file_format: str, file_path: str,
+                         cache_max_size: int):
+    file_path_for_pyarrow = file_io.to_filesystem_path(file_path)
+    filesystem = file_io.filesystem
+
+    def load():
+        return ds.dataset(
+            file_path_for_pyarrow, format=file_format, filesystem=filesystem)
+
+    key = (_FilesystemIdentity(filesystem), file_format, file_path_for_pyarrow)
+    if cache_max_size <= 0:
+        _reset_file_format_dataset_cache()
+        return load()
+
+    return _file_format_dataset_cache(cache_max_size).get_or_load(
+        key,
+        load,
+        lambda dataset: _estimate_file_format_cache_entry_size(
+            key, dataset, file_format))
+
+
 class FormatPyArrowReader(RecordBatchReader):
     """
     A Format Reader that reads record batch from a Parquet or ORC file using 
PyArrow,
@@ -63,7 +256,9 @@ class FormatPyArrowReader(RecordBatchReader):
                  row_ranges: Optional[List[Tuple[int, int]]] = None):
         self._predicate_field_names = predicate_field_names or set()
         file_path_for_pyarrow = file_io.to_filesystem_path(file_path)
-        self.dataset = ds.dataset(file_path_for_pyarrow, format=file_format, 
filesystem=file_io.filesystem)
+        cache_max_size = _file_format_metadata_cache_max_size(file_io)
+        self.dataset = _file_format_dataset(
+            file_io, file_format, file_path, cache_max_size)
         self._range_slicer = None
         self._selected_parquet_row_groups = None
         self._exhausted = False
diff --git a/paimon-python/pypaimon/tests/parquet_metadata_cache_test.py 
b/paimon-python/pypaimon/tests/parquet_metadata_cache_test.py
new file mode 100644
index 0000000000..2ca91f5776
--- /dev/null
+++ b/paimon-python/pypaimon/tests/parquet_metadata_cache_test.py
@@ -0,0 +1,479 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements.  See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership.  The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License.  You may obtain a copy of the License at
+#
+#     http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import gc
+import os
+import tempfile
+import threading
+import time
+import unittest
+import weakref
+from concurrent.futures import Future, ThreadPoolExecutor
+from unittest.mock import patch
+
+import pyarrow as pa
+import pyarrow.fs as pafs
+import pyarrow.parquet as pq
+from fsspec.implementations.local import LocalFileSystem as 
FsspecLocalFileSystem
+
+from pypaimon.common.options import Options
+from pypaimon.common.options.config import CatalogOptions
+from pypaimon.common.options.core_options import CoreOptions
+from pypaimon.filesystem.local_file_io import LocalFileIO
+from pypaimon.read.reader import format_pyarrow_reader as reader_module
+from pypaimon.read.reader.format_pyarrow_reader import FormatPyArrowReader
+from pypaimon.schema.data_types import AtomicType, DataField
+
+
+DEFAULT_CACHE_SIZE = 50 * 1024 * 1024
+
+
+class _CountingInputFile:
+    def __init__(self, wrapped, file_system):
+        self._wrapped = wrapped
+        self._file_system = file_system
+
+    def read(self, size=-1):
+        offset = self._wrapped.tell()
+        data = self._wrapped.read(size)
+        self._file_system.reads.append((offset, len(data)))
+        return data
+
+    def readinto(self, buffer):
+        offset = self._wrapped.tell()
+        size = self._wrapped.readinto(buffer)
+        self._file_system.reads.append((offset, size))
+        return size
+
+    def __getattr__(self, name):
+        return getattr(self._wrapped, name)
+
+
+class _CountingLocalFileSystem(FsspecLocalFileSystem):
+    def __init__(self):
+        super().__init__()
+        self.opens = 0
+        self.reads = []
+
+    def _open(self, path, mode="rb", **kwargs):
+        wrapped = super()._open(path, mode=mode, **kwargs)
+        if "r" not in mode:
+            return wrapped
+        self.opens += 1
+        return _CountingInputFile(wrapped, self)
+
+    def reset_counts(self):
+        self.opens = 0
+        self.reads = []
+
+
+class FileFormatMetadataCacheTest(unittest.TestCase):
+    def setUp(self):
+        reader_module._reset_file_format_dataset_cache()
+        self.temp_dir = tempfile.TemporaryDirectory()
+        self.file_io = LocalFileIO(self.temp_dir.name, Options({}))
+        self.paths = []
+        for index in range(3):
+            path = os.path.join(self.temp_dir.name, 
"data-{}.parquet".format(index))
+            pq.write_table(
+                pa.table({"value": list(range(index * 10, index * 10 + 10))}),
+                path,
+                row_group_size=2,
+            )
+            self.paths.append(path)
+
+    def tearDown(self):
+        reader_module._reset_file_format_dataset_cache()
+        self.temp_dir.cleanup()
+
+    def _file_io(self, max_size="50 mb"):
+        return LocalFileIO(self.temp_dir.name, Options({
+            "file-format.metadata-cache.max-size": max_size,
+        }))
+
+    def _read(self, path, file_io=None, options=None):
+        reader = FormatPyArrowReader(
+            file_io or self.file_io,
+            "parquet",
+            path,
+            [DataField(0, "value", AtomicType("BIGINT"))],
+            None,
+            options=options,
+        )
+        values = []
+        try:
+            while True:
+                batch = reader.read_arrow_batch()
+                if batch is None:
+                    return values
+                values.extend(batch.column(0).to_pylist())
+        finally:
+            reader.close()
+
+    def test_enabled_by_default(self):
+        self.assertEqual(
+            DEFAULT_CACHE_SIZE,
+            self.file_io.properties.get(
+                
CatalogOptions.FILE_FORMAT_METADATA_CACHE_MAX_SIZE).get_bytes())
+
+        original = reader_module.ds.dataset
+        with patch.object(reader_module.ds, "dataset", wraps=original) as 
dataset:
+            self._read(self.paths[0])
+            self._read(self.paths[0])
+        self.assertEqual(1, dataset.call_count)
+
+    def test_zero_size_bypasses_and_removes_entry(self):
+        enabled = self._file_io()
+        disabled = self._file_io("0 b")
+        original = reader_module.ds.dataset
+        with patch.object(reader_module.ds, "dataset", wraps=original) as 
dataset:
+            self._read(self.paths[0], enabled)
+            self._read(self.paths[0], disabled)
+            self._read(self.paths[0], enabled)
+        self.assertEqual(3, dataset.call_count)
+
+    def test_zero_size_clears_other_entries(self):
+        enabled = self._file_io()
+        disabled = self._file_io("0 b")
+        original = reader_module.ds.dataset
+        with patch.object(reader_module.ds, "dataset", wraps=original) as 
dataset:
+            self._read(self.paths[0], enabled)
+            self._read(self.paths[1], enabled)
+            self._read(self.paths[0], disabled)
+            self._read(self.paths[1], enabled)
+        self.assertEqual(4, dataset.call_count)
+
+    def test_reuses_dataset(self):
+        file_io = self._file_io()
+        original = reader_module.ds.dataset
+        with patch.object(reader_module.ds, "dataset", wraps=original) as 
dataset:
+            first = self._read(self.paths[0], file_io)
+            second = self._read(self.paths[0], file_io)
+
+        self.assertEqual(list(range(10)), first)
+        self.assertEqual(first, second)
+        self.assertEqual(1, dataset.call_count)
+
+    def test_repeated_scan_skips_footer_io(self):
+        path = os.path.join(self.temp_dir.name, "footer-io.parquet")
+        pq.write_table(
+            pa.table({
+                "value": list(range(10000)),
+                "payload": ["x" * 100] * 10000,
+            }),
+            path,
+            row_group_size=100,
+            compression="none",
+        )
+        counting = _CountingLocalFileSystem()
+        file_io = self._file_io("0 b")
+        file_io.filesystem = pafs.PyFileSystem(pafs.FSSpecHandler(counting))
+
+        uncached = self._read(path, file_io)
+        uncached_opens = counting.opens
+        uncached_reads = len(counting.reads)
+
+        counting.reset_counts()
+        reader_module._reset_file_format_dataset_cache()
+        file_io.properties.set(
+            CatalogOptions.FILE_FORMAT_METADATA_CACHE_MAX_SIZE, "50 mb")
+        self._read(path, file_io)
+        counting.reset_counts()
+        cached = self._read(path, file_io)
+
+        self.assertEqual(uncached, cached)
+        self.assertLess(counting.opens, uncached_opens)
+        self.assertLess(len(counting.reads), uncached_reads)
+
+    def test_evicts_least_recently_used_entry_by_estimated_size(self):
+        cache = reader_module._FileFormatDatasetCache(10)
+        first_key = (None, "parquet", "first")
+        second_key = (None, "parquet", "second")
+        third_key = (None, "parquet", "third")
+
+        cache.get_or_load(first_key, lambda: "first", lambda _: 4)
+        cache.get_or_load(second_key, lambda: "second", lambda _: 4)
+        cache.get_or_load(first_key, lambda: "unused", lambda _: 4)
+        cache.get_or_load(third_key, lambda: "third", lambda _: 4)
+
+        self.assertEqual([first_key, third_key], list(cache._entries.keys()))
+        self.assertEqual(8, cache.estimated_size)
+
+    def test_evicts_least_recently_used_entry_by_entry_count(self):
+        cache = reader_module._FileFormatDatasetCache(
+            1024, max_entries=2)
+        first_key = (None, "parquet", "first")
+        second_key = (None, "parquet", "second")
+        third_key = (None, "parquet", "third")
+
+        cache.get_or_load(first_key, lambda: "first", lambda _: 1)
+        cache.get_or_load(second_key, lambda: "second", lambda _: 1)
+        cache.get_or_load(third_key, lambda: "third", lambda _: 1)
+
+        self.assertEqual(
+            [second_key, third_key], list(cache._entries.keys()))
+        self.assertEqual(2, cache.estimated_size)
+
+    def test_does_not_retain_entry_larger_than_size_limit(self):
+        cache = reader_module._FileFormatDatasetCache(5)
+        loads = []
+        key = (None, "parquet", "large")
+
+        def load():
+            loads.append(True)
+            return "large"
+
+        self.assertEqual(
+            "large", cache.get_or_load(key, load, lambda _: 6))
+        self.assertEqual(
+            "large", cache.get_or_load(key, load, lambda _: 6))
+        self.assertEqual(2, len(loads))
+        self.assertEqual(0, len(cache._entries))
+        self.assertEqual(0, cache.estimated_size)
+
+    def test_does_not_retain_entry_without_size_estimate(self):
+        cache = reader_module._FileFormatDatasetCache(10)
+        key = (None, "unknown", "data")
+        loads = []
+
+        def load():
+            loads.append(True)
+            return "unknown"
+
+        self.assertEqual(
+            "unknown", cache.get_or_load(key, load, lambda _: None))
+        self.assertEqual(
+            "unknown", cache.get_or_load(key, load, lambda _: None))
+        self.assertEqual(2, len(loads))
+        self.assertEqual(0, len(cache._entries))
+
+    def test_coalesces_load_while_uncached_result_completes(self):
+        cache = reader_module._FileFormatDatasetCache(10)
+        key = (None, "unknown", "data")
+        setting_result = threading.Event()
+        waiter_started = threading.Event()
+        release_result = threading.Event()
+        loads = []
+        future_count = []
+
+        def new_future():
+            future = Future()
+            future_count.append(future)
+            if len(future_count) == 1:
+                original_set_result = future.set_result
+                original_result = future.result
+
+                def delayed_set_result(result):
+                    setting_result.set()
+                    release_result.wait()
+                    original_set_result(result)
+
+                def observed_result(*args, **kwargs):
+                    waiter_started.set()
+                    return original_result(*args, **kwargs)
+
+                future.set_result = delayed_set_result
+                future.result = observed_result
+            return future
+
+        def load():
+            loads.append(True)
+            return "unknown"
+
+        with patch.object(reader_module, "Future", side_effect=new_future):
+            with ThreadPoolExecutor(max_workers=2) as executor:
+                first = executor.submit(
+                    cache.get_or_load, key, load, lambda _: None)
+                self.assertTrue(setting_result.wait(1))
+                second = executor.submit(
+                    cache.get_or_load, key, load, lambda _: None)
+                waited = waiter_started.wait(1)
+                release_result.set()
+
+                self.assertTrue(waited)
+                self.assertEqual("unknown", first.result())
+                self.assertEqual("unknown", second.result())
+        self.assertEqual(1, len(loads))
+
+    def test_estimates_serialized_parquet_footer_size(self):
+        dataset = reader_module.ds.dataset(self.paths[0], format="parquet")
+        expected = sum(
+            fragment.metadata.serialized_size
+            for fragment in dataset.get_fragments()
+        )
+        self.assertGreater(expected, 0)
+        self.assertEqual(
+            expected,
+            reader_module._estimate_file_format_dataset_size(
+                dataset, "parquet"))
+
+    def test_cache_entry_estimate_has_conservative_floor(self):
+        dataset = reader_module.ds.dataset(
+            self.paths[0], format="parquet")
+        key = (
+            reader_module._FilesystemIdentity(self.file_io.filesystem),
+            "parquet",
+            self.paths[0],
+        )
+
+        estimated = reader_module._estimate_file_format_cache_entry_size(
+            key, dataset, "parquet")
+
+        self.assertGreaterEqual(
+            estimated,
+            reader_module._FILE_FORMAT_METADATA_CACHE_MIN_ENTRY_SIZE)
+
+    def test_process_cache_can_shrink_requested_capacity(self):
+        cache = reader_module._file_format_dataset_cache(10)
+        cache.get_or_load(("first", "parquet", "first"),
+                          lambda: "first", lambda _: 4)
+        cache.get_or_load(("second", "parquet", "second"),
+                          lambda: "second", lambda _: 4)
+        same_cache = reader_module._file_format_dataset_cache(5)
+
+        self.assertIs(cache, same_cache)
+        self.assertEqual(5, cache.max_size)
+        self.assertEqual(
+            [("second", "parquet", "second")],
+            list(cache._entries.keys()))
+        self.assertEqual(4, cache.estimated_size)
+
+    def test_table_option_does_not_configure_process_cache(self):
+        table_options = CoreOptions(Options({
+            "file-format.metadata-cache.max-size": "0 b",
+        }))
+        original = reader_module.ds.dataset
+        with patch.object(reader_module.ds, "dataset", wraps=original) as 
dataset:
+            self._read(self.paths[0], options=table_options)
+            self._read(self.paths[0], options=table_options)
+        self.assertEqual(1, dataset.call_count)
+
+    def test_shares_cache_across_file_io_with_same_filesystem(self):
+        other_file_io = LocalFileIO(self.temp_dir.name, Options({}))
+        other_file_io.filesystem = self.file_io.filesystem
+
+        original = reader_module.ds.dataset
+        with patch.object(reader_module.ds, "dataset", wraps=original) as 
dataset:
+            reader_module._file_format_dataset(
+                self.file_io, "parquet", self.paths[0], DEFAULT_CACHE_SIZE)
+            reader_module._file_format_dataset(
+                other_file_io, "parquet", self.paths[0], DEFAULT_CACHE_SIZE)
+
+        self.assertEqual(1, dataset.call_count)
+
+    def test_does_not_share_across_filesystems(self):
+        other_file_io = LocalFileIO(self.temp_dir.name, Options({}))
+        original = reader_module.ds.dataset
+        with patch.object(reader_module.ds, "dataset", wraps=original) as 
dataset:
+            reader_module._file_format_dataset(
+                self.file_io, "parquet", self.paths[0], DEFAULT_CACHE_SIZE)
+            reader_module._file_format_dataset(
+                other_file_io, "parquet", self.paths[0], DEFAULT_CACHE_SIZE)
+        self.assertEqual(2, dataset.call_count)
+
+    def test_does_not_share_across_file_formats(self):
+        parquet_dataset = object()
+        orc_dataset = object()
+        with patch.object(
+                reader_module.ds, "dataset",
+                side_effect=[parquet_dataset, orc_dataset]) as dataset:
+            with patch.object(
+                    reader_module, "_estimate_file_format_dataset_size",
+                    return_value=1):
+                first = reader_module._file_format_dataset(
+                    self.file_io, "parquet", self.paths[0], DEFAULT_CACHE_SIZE)
+                second = reader_module._file_format_dataset(
+                    self.file_io, "orc", self.paths[0], DEFAULT_CACHE_SIZE)
+
+        self.assertIs(parquet_dataset, first)
+        self.assertIs(orc_dataset, second)
+        self.assertEqual(2, dataset.call_count)
+
+    def test_cache_key_retains_filesystem_wrapper(self):
+        root = pafs.LocalFileSystem()
+        filesystem = pafs.SubTreeFileSystem(self.temp_dir.name, root)
+        filesystem_ref = weakref.ref(filesystem)
+        file_io = LocalFileIO(self.temp_dir.name, Options({}))
+        file_io.filesystem = filesystem
+
+        reader_module._file_format_dataset(
+            file_io, "parquet", os.path.basename(self.paths[0]),
+            DEFAULT_CACHE_SIZE)
+        file_io.filesystem = root
+        del filesystem
+        gc.collect()
+
+        self.assertIsNotNone(filesystem_ref())
+
+    def test_filesystem_hash_collision_does_not_share_dataset(self):
+        first_dir = tempfile.TemporaryDirectory()
+        second_dir = tempfile.TemporaryDirectory()
+        try:
+            file_name = "same.parquet"
+            pq.write_table(
+                pa.table({"value": [1]}), os.path.join(first_dir.name, 
file_name))
+            pq.write_table(
+                pa.table({"value": [2]}), os.path.join(second_dir.name, 
file_name))
+            file_io = LocalFileIO(first_dir.name, Options({}))
+            first_filesystem = pafs.SubTreeFileSystem(
+                first_dir.name, pafs.LocalFileSystem())
+            second_filesystem = pafs.SubTreeFileSystem(
+                second_dir.name, pafs.LocalFileSystem())
+
+            with patch.object(
+                    reader_module._FilesystemIdentity, "__hash__", 
return_value=1):
+                file_io.filesystem = first_filesystem
+                first = reader_module._file_format_dataset(
+                    file_io, "parquet", file_name,
+                    DEFAULT_CACHE_SIZE).to_table()
+                file_io.filesystem = second_filesystem
+                second = reader_module._file_format_dataset(
+                    file_io, "parquet", file_name,
+                    DEFAULT_CACHE_SIZE).to_table()
+
+            self.assertEqual([1], first.column("value").to_pylist())
+            self.assertEqual([2], second.column("value").to_pylist())
+        finally:
+            first_dir.cleanup()
+            second_dir.cleanup()
+
+    def test_resets_after_process_change(self):
+        parent_cache = 
reader_module._file_format_dataset_cache(DEFAULT_CACHE_SIZE)
+        with patch.object(reader_module.os, "getpid", return_value=os.getpid() 
+ 1):
+            child_cache = 
reader_module._file_format_dataset_cache(DEFAULT_CACHE_SIZE)
+        self.assertIsNot(parent_cache, child_cache)
+
+    def test_coalesces_concurrent_loads(self):
+        original = reader_module.ds.dataset
+
+        def delayed_dataset(*args, **kwargs):
+            time.sleep(0.05)
+            return original(*args, **kwargs)
+
+        with patch.object(
+                reader_module.ds, "dataset", side_effect=delayed_dataset) as 
dataset:
+            with ThreadPoolExecutor(max_workers=8) as executor:
+                results = list(executor.map(
+                    lambda _: self._read(self.paths[0]),
+                    range(8),
+                ))
+
+        self.assertEqual(1, dataset.call_count)
+        self.assertTrue(all(value == list(range(10)) for value in results))
+
+
+if __name__ == "__main__":
+    unittest.main()

Reply via email to