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 42e77ea3a0 [python] Add as-of joins for multimodal scans (#9536)
42e77ea3a0 is described below

commit 42e77ea3a06e7438c99e9666ac9b6bcee49f19a9
Author: XiaoHongbo <[email protected]>
AuthorDate: Fri Sep 11 17:55:26 2026 +0800

    [python] Add as-of joins for multimodal scans (#9536)
---
 docs/docs/pypaimon/multimodal-reading.md           |   39 +
 paimon-python/pypaimon/multimodal/__init__.py      |    6 +
 paimon-python/pypaimon/multimodal/temporal.py      | 1025 ++++++++++++++
 .../pypaimon/read/reader/format_pyarrow_reader.py  |   90 +-
 paimon-python/pypaimon/read/split_read.py          |    4 +-
 paimon-python/pypaimon/read/table_read.py          |    2 +
 .../pypaimon/tests/multimodal_temporal_test.py     | 1411 ++++++++++++++++++++
 .../pypaimon/tests/parquet_metadata_cache_test.py  |    6 +-
 .../pypaimon/tests/parquet_row_range_test.py       |   54 +-
 9 files changed, 2623 insertions(+), 14 deletions(-)

diff --git a/docs/docs/pypaimon/multimodal-reading.md 
b/docs/docs/pypaimon/multimodal-reading.md
index 35ffff8f12..2611154413 100644
--- a/docs/docs/pypaimon/multimodal-reading.md
+++ b/docs/docs/pypaimon/multimodal-reading.md
@@ -50,6 +50,45 @@ with docs.scan().where("category = 
'lake'").to_arrow_batch_reader() as reader:
         consume(batch)
 ```
 
+### As-of joins
+
+`join_asof` preserves each left row and matches at most one right row in the
+same `by` group. Chain calls to align multiple streams lazily.
+
+```python
+from datetime import timedelta
+from pypaimon.multimodal import join_asof
+
+aligned = join_asof(
+    actions.scan().select(["episode_id", "event_time", "action"]),
+    images.scan().where("camera = 'left'").select("image"),
+    on="event_time",
+    by="episode_id",
+    direction="nearest",
+    tolerance=timedelta(milliseconds=20),
+).join_asof(
+    topics.scan().where("topic = '/robot/state'").select("value"),
+    direction="backward",
+    tolerance=timedelta(milliseconds=50),
+)
+
+for batch in aligned.to_arrow_batch_reader(batch_size=128):
+    train(batch)
+```
+
+`direction` is `backward`, `forward`, or `nearest`; tolerance is inclusive and
+zero means exact. Nearest ties use the earlier time. For duplicate timestamps,
+backward uses the last row and forward uses the first. Nearest uses the last
+row for an exact match; otherwise it uses the backward or forward candidate's
+rule. Misses return null.
+
+Keys must be non-null with matching types. Use `right_on` for a different right
+timestamp and `suffix` for conflicts. Select the right timestamp to compute the
+match delta.
+
+Inputs are snapshot-pinned (`resolved_snapshots`). Left rows stream, right join
+keys stay in memory, and BLOBs remain descriptors.
+
 ### Reading BLOB columns
 
 `scan().read_blobs(column)` bulk-fetches a BLOB column's bytes for the filtered
diff --git a/paimon-python/pypaimon/multimodal/__init__.py 
b/paimon-python/pypaimon/multimodal/__init__.py
index 584edc8431..eec6b9afb6 100644
--- a/paimon-python/pypaimon/multimodal/__init__.py
+++ b/paimon-python/pypaimon/multimodal/__init__.py
@@ -42,6 +42,10 @@ from pypaimon.multimodal.table import (
     text_route,
     vector_route,
 )
+from pypaimon.multimodal.temporal import (
+    AsOfJoin,
+    join_asof,
+)
 from pypaimon.multimodal.video import VideoFrameCollator
 from pypaimon.table.row.blob import Blob, BlobDescriptor, VideoFrameDescriptor
 from pypaimon.table.data_evolution_merge_into import (
@@ -51,6 +55,7 @@ from pypaimon.table.data_evolution_merge_into import (
 )
 
 __all__ = [
+    "AsOfJoin",
     "Blob",
     "BlobDescriptor",
     "BlobObject",
@@ -71,6 +76,7 @@ __all__ = [
     "VideoFrameCollator",
     "VideoFrameDescriptor",
     "connect",
+    "join_asof",
     "lit",
     "source_col",
     "target_col",
diff --git a/paimon-python/pypaimon/multimodal/temporal.py 
b/paimon-python/pypaimon/multimodal/temporal.py
new file mode 100644
index 0000000000..808e0fee53
--- /dev/null
+++ b/paimon-python/pypaimon/multimodal/temporal.py
@@ -0,0 +1,1025 @@
+# 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.
+
+"""Temporal alignment for multimodal table scans."""
+
+from bisect import bisect_left, bisect_right
+from datetime import timedelta
+import json
+import math
+from numbers import Integral, Real
+
+import pyarrow as pa
+import pyarrow.compute as pc
+
+from pypaimon.catalog.table_query_auth import TableQueryAuthResult
+from pypaimon.common.options.core_options import CoreOptions, StartupMode
+from pypaimon.common.predicate_json_parser import (
+    _apply_predicate_transform,
+    _collect_all_field_refs_from_transform,
+)
+from pypaimon.globalindex.indexed_split import IndexedSplit
+from pypaimon.multimodal.query import ScanQuery
+from pypaimon.read.query_auth_split import QueryAuthSplit, resolve_auth_result
+from pypaimon.read.reader.format_pyarrow_reader import _DecodedRowGroupCache
+from pypaimon.read.table_read import _ClosableArrowBatchReader
+from pypaimon.schema.data_types import PyarrowFieldParser
+from pypaimon.snapshot.time_travel_util import TimeTravelUtil
+from pypaimon.table.special_fields import SpecialFields
+from pypaimon.table.source.global_index_live_row_filter import (
+    table_at_snapshot,
+)
+from pypaimon.utils.range import Range
+
+
+_ROW_ID = SpecialFields.ROW_ID.name
+_MAX_INT64 = (1 << 63) - 1
+_TIME_KEY = object()
+_TEMPORAL_ROW_GROUP_CACHE_MAX_SIZE = 64 * 1024 * 1024
+
+
+def join_asof(left, right, *, on, by, direction="backward", tolerance=None,
+              right_on=None, suffix="_right") -> "AsOfJoin":
+    """Join each left row with at most one time-aligned right row."""
+    if not isinstance(on, str) or not on:
+        raise ValueError("on must be a non-empty column name.")
+    if isinstance(by, str):
+        by = (by,)
+    else:
+        try:
+            by = tuple(by)
+        except TypeError as error:
+            raise ValueError(
+                "by must be a column name or sequence.") from error
+    if not by:
+        raise ValueError(
+            "join_asof requires at least one grouping column in by.")
+    if (any(not isinstance(name, str) or not name for name in by)
+            or len(set(by)) != len(by)):
+        raise ValueError("by must contain unique, non-empty column names.")
+    return AsOfJoin(left, on, by).join_asof(
+        right,
+        direction=direction,
+        tolerance=tolerance,
+        right_on=right_on,
+        suffix=suffix,
+    )
+
+
+class AsOfJoin:
+    """Lazy, chainable result of :func:`join_asof`."""
+
+    def __init__(self, left, on, by):
+        self._anchor = _pin_scan_to_snapshot(_require_scan(left, "left"))
+        self._on = on
+        self._by = by
+        self._sources = ()
+        self._anchor_schema = _query_schema(self._anchor)
+        self._anchor_table_schema = _table_schema(self._anchor)
+        self._validate_anchor()
+        self.schema = self._output_schema()
+
+    def join_asof(self, right, *, direction="backward", tolerance=None,
+                  right_on=None, suffix="_right") -> "AsOfJoin":
+        """Append a right-side as-of join without materializing this scan."""
+        position = len(self._sources) + 1
+        label = "right source %d" % position
+        source = _AsOfJoinRight(
+            label,
+            right,
+            self._on,
+            self._by,
+            direction,
+            tolerance,
+            right_on,
+            suffix,
+        )
+
+        result = object.__new__(AsOfJoin)
+        result._anchor = self._anchor
+        result._on = self._on
+        result._by = self._by
+        result._sources = self._sources + (source,)
+        result._anchor_schema = self._anchor_schema
+        result._anchor_table_schema = self._anchor_table_schema
+        result._validate_anchor()
+        result.schema = result._output_schema()
+        return result
+
+    def to_arrow_batch_reader(self, *, batch_size=1024):
+        """Index right-side timestamps, then stream aligned rows in batches."""
+        if (isinstance(batch_size, bool)
+                or not isinstance(batch_size, int)
+                or batch_size <= 0):
+            raise ValueError("batch_size must be a positive integer.")
+
+        anchor_metadata = _metadata_batches(
+            self._anchor, self._on, self._by, batch_size)
+        row_group_cache = _DecodedRowGroupCache(
+            _TEMPORAL_ROW_GROUP_CACHE_MAX_SIZE)
+        anchor_fetcher = _RowIdFetcher(self._anchor, row_group_cache)
+        source_fetchers = []
+        for source in self._sources:
+            source.plan()
+            source_fetchers.append(
+                _RowIdFetcher(source.query, row_group_cache))
+        schema = self._output_schema(anchor_fetcher.schema, source_fetchers)
+        self.schema = schema
+
+        def batches():
+            try:
+                for metadata in anchor_metadata:
+                    rows = _metadata_rows(
+                        metadata, self._on,
+                        self._anchor_table_schema.field(self._on).type)
+                    yield from self._build_batches(
+                        rows, anchor_fetcher, source_fetchers, schema)
+            finally:
+                anchor_metadata.close()
+
+        batch_iterator = batches()
+        reader = pa.ipc.RecordBatchReader.from_batches(schema, batch_iterator)
+        return _ClosableArrowBatchReader(reader, batch_iterator)
+
+    def to_arrow(self):
+        reader = self.to_arrow_batch_reader()
+        try:
+            return reader.read_all()
+        finally:
+            close = getattr(reader, "close", None)
+            if close is not None:
+                close()
+
+    def to_pandas(self):
+        return self.to_arrow().to_pandas()
+
+    def to_list(self):
+        return _arrow_rows(self.to_arrow())
+
+    @property
+    def resolved_snapshots(self):
+        """Return the table snapshots pinned by this alignment."""
+        snapshots = {
+            "left": _resolved_snapshot(self._anchor),
+        }
+        snapshots.update({
+            "right_%d" % position: _resolved_snapshot(source.query)
+            for position, source in enumerate(self._sources, 1)
+        })
+        return snapshots
+
+    def _validate_anchor(self):
+        _require_columns(
+            self._anchor_table_schema, self._by + (self._on,), "anchor")
+        anchor_type = self._anchor_table_schema.field(self._on).type
+        _delta_type(anchor_type)
+        for name in self._by:
+            _validate_group_type(
+                name, self._anchor_table_schema.field(name).type)
+        for source in self._sources:
+            if source.time_type != anchor_type:
+                raise TypeError(
+                    "Left and %s temporal columns must have the same "
+                    "type; got %s and %s."
+                    % (source.label, anchor_type, source.time_type)
+                )
+            for name in self._by:
+                anchor_group_type = self._anchor_table_schema.field(name).type
+                source_group_type = source.table_schema.field(name).type
+                if source_group_type != anchor_group_type:
+                    raise TypeError(
+                        "Left and %s grouping column %r must have "
+                        "the same type; got %s and %s."
+                        % (source.label, name, anchor_group_type,
+                           source_group_type)
+                    )
+
+    def _output_schema(self, anchor_schema=None, source_fetchers=None):
+        if anchor_schema is None:
+            anchor_schema = self._anchor_schema
+        fields = list(anchor_schema)
+        names = set(anchor_schema.names)
+        for position, source in enumerate(self._sources):
+            payload_schema = (
+                source.payload_schema if source_fetchers is None
+                else source_fetchers[position].schema
+            )
+            for name in source.payload_schema.names:
+                field = payload_schema.field(name)
+                output_name = field.name
+                if output_name in names:
+                    output_name += source.suffix
+                if output_name in names:
+                    raise ValueError(
+                        "%s column %r conflicts after applying suffix %r."
+                        % (source.label, field.name, source.suffix)
+                    )
+                output = pa.field(
+                    output_name, field.type, nullable=True,
+                    metadata=field.metadata)
+                fields.append(output)
+                names.add(output.name)
+        return pa.schema(fields, metadata=anchor_schema.metadata)
+
+    def _build_batch(
+            self, anchor_rows, anchor_fetcher, source_fetchers, schema):
+        anchor_ids = [row[_ROW_ID] for row in anchor_rows]
+        anchor = anchor_fetcher.fetch(anchor_ids)
+        anchor.validate()
+        arrays = [anchor[name] for name in self._anchor_schema.names]
+
+        for source, fetcher in zip(self._sources, source_fetchers):
+            matches = [source.match(row) for row in anchor_rows]
+            matched_ids = [match for match in matches if match is not None]
+            unique_ids = list(dict.fromkeys(matched_ids))
+            values = fetcher.fetch(unique_ids)
+            positions = {
+                row_id: index for index, row_id in enumerate(unique_ids)
+            }
+            take = pa.array([
+                None if match is None else positions[match]
+                for match in matches
+            ], type=pa.int64())
+            for field in source.payload_schema:
+                array = pc.take(values[field.name], take)
+                array.validate()
+                arrays.append(array)
+
+        if arrays:
+            table = pa.Table.from_arrays(
+                arrays, schema=schema).combine_chunks()
+            return table.to_batches(max_chunksize=len(anchor_rows))[0]
+        batch = pa.RecordBatch.from_struct_array(pa.array(
+            [{}] * len(anchor_rows), type=pa.struct([])))
+        return batch.replace_schema_metadata(schema.metadata)
+
+    def _build_batches(
+            self, anchor_rows, anchor_fetcher, source_fetchers, schema):
+        try:
+            yield self._build_batch(
+                anchor_rows, anchor_fetcher, source_fetchers, schema)
+        except pa.ArrowInvalid as error:
+            if len(anchor_rows) < 2 or "offset" not in str(error).lower():
+                raise
+            middle = len(anchor_rows) // 2
+            yield from self._build_batches(
+                anchor_rows[:middle], anchor_fetcher,
+                source_fetchers, schema)
+            yield from self._build_batches(
+                anchor_rows[middle:], anchor_fetcher,
+                source_fetchers, schema)
+
+
+class _AsOfJoinRight:
+
+    def __init__(self, label, query, anchor_on, by, direction, tolerance,
+                 right_on, suffix):
+        _validate_join_options(direction, tolerance, right_on, suffix)
+        self.label = label
+        self.query = _pin_scan_to_snapshot(_require_scan(query, label))
+        self.direction = direction
+        self.suffix = suffix
+        self.anchor_on = anchor_on
+        self.on = anchor_on if right_on is None else right_on
+        self.by = by
+        self.table_schema = _table_schema(self.query)
+        _require_columns(
+            self.table_schema, by + (self.on,), label)
+        self.time_type = self.table_schema.field(self.on).type
+        _delta_type(self.time_type)
+        _validate_tolerance(tolerance, self.time_type)
+        self._tolerance_key = _time_tolerance_key(tolerance, self.time_type)
+        schema, paths = _query_schema_and_paths(self.query)
+        projection = self.query._effective_projection()
+        excluded = {(name,) for name in by}
+        if projection is None:
+            excluded.add((self.on,))
+        self.payload_schema = pa.schema([
+            field for field, path in zip(schema, paths)
+            if tuple(path) not in excluded
+        ])
+        self._index = None
+
+    def plan(self):
+        metadata = _metadata_table(self.query, self.on, self.by)
+        self._times = metadata[self.on].combine_chunks()
+        self._time_keys = _time_search_keys(self._times, self.time_type)
+        self._row_ids = metadata[_ROW_ID].combine_chunks()
+        self._index = {}
+        group_columns = [metadata[name].combine_chunks() for name in self.by]
+        previous = None
+        start = 0
+        for position in range(len(metadata)):
+            key = tuple(column[position].as_py() for column in group_columns)
+            if position and key != previous:
+                self._index[previous] = (start, position)
+                start = position
+            previous = key
+        if len(metadata):
+            self._index[previous] = (start, len(metadata))
+
+    def match(self, anchor_row):
+        key = tuple(anchor_row[name] for name in self.by)
+        bounds = self._index.get(key)
+        if bounds is None:
+            return None
+        target_key = anchor_row[_TIME_KEY]
+        index = _match_index(
+            self._time_keys, target_key, self.direction, *bounds)
+        if index is None:
+            return None
+        matched_key = _python_scalar(self._time_keys[index])
+        if (self._tolerance_key is not None
+                and abs(matched_key - target_key) > self._tolerance_key):
+            return None
+        return self._row_ids[index].as_py()
+
+
+def _validate_join_options(direction, tolerance, right_on, suffix):
+    if direction not in ("backward", "forward", "nearest"):
+        raise ValueError(
+            "direction must be 'backward', 'forward', or 'nearest'.")
+    if right_on is not None and (
+            not isinstance(right_on, str) or not right_on):
+        raise ValueError("right_on must be a non-empty column name.")
+    if not isinstance(suffix, str):
+        raise TypeError("suffix must be a string.")
+    if tolerance is not None:
+        if isinstance(tolerance, bool) or not isinstance(
+                tolerance, (Real, timedelta)):
+            raise TypeError("tolerance must be numeric or datetime.timedelta.")
+        if (isinstance(tolerance, Real)
+                and not isinstance(tolerance, Integral)
+                and not math.isfinite(tolerance)):
+            raise ValueError("tolerance must be finite.")
+        zero = timedelta(0) if isinstance(tolerance, timedelta) else 0
+        if tolerance < zero:
+            raise ValueError("tolerance must be non-negative.")
+
+
+def _require_scan(query, label):
+    if (type(query) is not ScanQuery
+            or getattr(query, "_result_factory", None) is not None):
+        raise TypeError("%s must be a MultimodalTable.scan() query." % label)
+    return query
+
+
+def _pin_scan_to_snapshot(query):
+    table = query._table
+    options = table.options
+    if not options.row_tracking_enabled(False):
+        raise ValueError(
+            "join_asof requires 'row-tracking.enabled' = 'true'.")
+    if (options.scan_mode() == StartupMode.INCREMENTAL
+            or options.options.contains(
+                CoreOptions.INCREMENTAL_BETWEEN_TIMESTAMP)):
+        raise ValueError(
+            "join_asof does not support incremental scans; inputs must "
+            "represent a complete point-in-time snapshot.")
+    # Validate the original scan configuration before replacing it with a
+    # pinned snapshot. Otherwise an invalid or unsupported scan mode can be
+    # silently converted into a latest-full scan.
+    table.new_read_builder().new_scan()
+    snapshot = TimeTravelUtil.try_travel_to_snapshot(
+        options.options, table.tag_manager(), table.snapshot_manager())
+    if snapshot is None:
+        snapshot = table.snapshot_manager().get_latest_snapshot()
+    empty = snapshot is None
+    tag_name = (
+        options.scan_tag_name()
+        if options.options.contains_key(CoreOptions.SCAN_TAG_NAME.key())
+        else None
+    )
+    if snapshot is not None and tag_name is None:
+        table = table_at_snapshot(table, snapshot)
+    pinned = ScanQuery(table)
+    pinned._predicate = query._predicate
+    pinned._projection = query._projection
+    pinned._limit = query._limit
+    pinned._include_row_id = query._include_row_id
+    pinned._temporal_empty = empty
+    pinned._temporal_snapshot_id = (
+        None if snapshot is None else snapshot.id)
+    pinned._temporal_tag_name = tag_name
+    return pinned
+
+
+def _resolved_snapshot(query):
+    resolved = {
+        "table": query._table.identifier.get_full_name(),
+        "snapshot_id": query._temporal_snapshot_id,
+    }
+    tag_name = getattr(query, "_temporal_tag_name", None)
+    if tag_name is not None:
+        resolved["tag_name"] = tag_name
+    return resolved
+
+
+def _query_schema(query):
+    return _query_schema_and_paths(query)[0]
+
+
+def _query_schema_and_paths(query):
+    table = query._table.copy_without_time_travel({
+        CoreOptions.BLOB_AS_DESCRIPTOR.key(): "true",
+    })
+    builder = query._configured_read_builder(table)
+    schema = PyarrowFieldParser.from_paimon_schema(builder.read_type())
+    paths = builder._nested_name_paths()
+    if paths is None:
+        paths = [[field.name] for field in schema]
+    return schema, paths
+
+
+def _table_schema(query):
+    return PyarrowFieldParser.from_paimon_schema(query._table.fields)
+
+
+def _metadata_table(query, on, by):
+    read_builder, splits, key_columns, output_columns = (
+        _metadata_builders(query, on, by))
+    if getattr(query, "_temporal_empty", False):
+        return _empty_metadata(query, key_columns, output_columns)
+    arrow = read_builder.new_read().to_arrow(splits)
+    metadata = arrow.select(output_columns).combine_chunks()
+    _validate_metadata(query, metadata, key_columns)
+    sort_keys = [(name, "ascending") for name in output_columns]
+    return metadata.take(pc.sort_indices(metadata, sort_keys=sort_keys))
+
+
+def _metadata_batches(query, on, by, batch_size):
+    read_builder, splits, key_columns, output_columns = (
+        _metadata_builders(query, on, by))
+    if getattr(query, "_temporal_empty", False):
+        return
+    reader = read_builder.new_read()._to_managed_arrow_batch_reader(splits)
+    try:
+        for batch in reader:
+            metadata = pa.Table.from_batches([batch]).select(output_columns)
+            _validate_metadata(query, metadata, key_columns)
+            for start in range(0, len(metadata), batch_size):
+                yield metadata.slice(start, batch_size)
+    finally:
+        reader.close()
+
+
+def _metadata_builders(query, on, by):
+    _validate_pinned_tag(query)
+    key_columns = list(dict.fromkeys(by + (on,)))
+    output_columns = key_columns + [_ROW_ID]
+    table = query._table.copy_without_time_travel({
+        CoreOptions.BLOB_AS_DESCRIPTOR.key(): "true",
+    })
+    plan_builder = table.new_read_builder()
+    if query._predicate is not None:
+        plan_builder = plan_builder.with_filter(query._predicate)
+    plan_builder = plan_builder.with_projection(key_columns)
+    if query._limit is not None:
+        plan_builder = plan_builder.with_limit(query._limit)
+    splits, masking = _plan_with_internal_row_id(
+        plan_builder,
+        query._temporal_snapshot_id,
+        getattr(query, "_temporal_empty", False),
+    )
+    key_masking = {
+        name: masking[name] for name in key_columns if name in masking
+    }
+    dependencies = _mask_dependencies(
+        key_masking, key_columns, _table_schema(query))
+    splits = _with_active_masking(splits, key_columns)
+    read_projection = list(dict.fromkeys(
+        output_columns + dependencies))
+    read_builder = table.new_read_builder().with_projection(read_projection)
+    if query._predicate is not None:
+        read_builder = read_builder.with_filter(query._predicate)
+    if query._limit is not None:
+        read_builder = read_builder.with_limit(query._limit)
+
+    physical_schema = PyarrowFieldParser.from_paimon_schema(
+        read_builder.read_type())
+    effective_schema = _effective_masked_schema(
+        physical_schema, key_masking)
+    for name in key_columns:
+        physical_type = physical_schema.field(name).type
+        effective_type = effective_schema.field(name).type
+        if effective_type != physical_type:
+            raise TypeError(
+                "Temporal key %r must preserve its type after column "
+                "masking; got %s instead of %s."
+                % (name, effective_type, physical_type)
+            )
+    return read_builder, splits, key_columns, output_columns
+
+
+def _empty_metadata(query, key_columns, output_columns):
+    return pa.Table.from_arrays([
+        pa.array([], type=_table_schema(query).field(name).type)
+        for name in key_columns
+    ] + [pa.array([], type=pa.int64())], names=output_columns)
+
+
+def _validate_metadata(query, metadata, key_columns):
+    for name in key_columns:
+        column = metadata[name]
+        if column.null_count:
+            raise ValueError("Temporal key %r cannot be null." % name)
+        if pa.types.is_floating(column.type):
+            for scalar in column:
+                if not math.isfinite(scalar.as_py()):
+                    raise ValueError(
+                        "Temporal key %r must be finite." % name)
+
+
+class _RowIdFetcher:
+
+    def __init__(self, query, row_group_cache):
+        _validate_pinned_tag(query)
+        self._schema = _query_schema(query)
+        table = query._table.copy_without_time_travel({
+            CoreOptions.BLOB_AS_DESCRIPTOR.key(): "true",
+        })
+        visible_projection = query._effective_projection()
+        plan_builder = table.new_read_builder()
+        if visible_projection is not None:
+            plan_projection = visible_projection
+            projected = table.new_read_builder().with_projection(
+                visible_projection)
+            projected_paths = projected._nested_name_paths()
+            if projected_paths is not None:
+                plan_projection = list(dict.fromkeys(
+                    path[0] for path in projected_paths))
+            plan_builder = plan_builder.with_projection(plan_projection)
+        if query._predicate is not None:
+            plan_builder = plan_builder.with_filter(query._predicate)
+        self._splits, masking = _plan_with_internal_row_id(
+            plan_builder,
+            query._temporal_snapshot_id,
+            getattr(query, "_temporal_empty", False),
+        )
+
+        read_projection = (
+            [field.name for field in table.fields]
+            if not visible_projection else list(visible_projection)
+        )
+        projected_builder = table.new_read_builder().with_projection(
+            list(dict.fromkeys(read_projection + [_ROW_ID])))
+        projected_schema = PyarrowFieldParser.from_paimon_schema(
+            projected_builder.read_type())
+        projected_paths = projected_builder._nested_name_paths()
+        if projected_paths is not None:
+            for field, path in zip(projected_schema, projected_paths):
+                if field.name in masking and field.name != path[0]:
+                    raise ValueError(
+                        "Temporal alignment cannot safely apply column "
+                        "masking to nested projection %r."
+                        % ".".join(path)
+                    )
+        if projected_paths is None:
+            active_targets = projected_schema.names
+        else:
+            active_targets = list(dict.fromkeys(
+                path[0] for path in projected_paths))
+        dependencies = _mask_dependencies(
+            masking, active_targets, _table_schema(query))
+        read_projection = list(dict.fromkeys(
+            read_projection + dependencies + [_ROW_ID]))
+        visible_builder = table.new_read_builder().with_projection(
+            read_projection)
+        self._fetch_schema = PyarrowFieldParser.from_paimon_schema(
+            visible_builder.read_type())
+        self._name_paths = visible_builder._nested_name_paths()
+        self._row_id_name = _ROW_ID
+        if self._name_paths is not None:
+            self._row_id_name = next(
+                field.name
+                for field, path in zip(
+                    self._fetch_schema, self._name_paths)
+                if path == [_ROW_ID]
+            )
+        if self._name_paths is None:
+            builder = visible_builder
+        else:
+            top_level = list(dict.fromkeys(
+                path[0] for path in self._name_paths))
+            builder = table.new_read_builder().with_projection(top_level)
+        if query._predicate is not None:
+            builder = builder.with_filter(query._predicate)
+        self._read = builder.new_read()
+        self._read._parquet_row_group_cache = row_group_cache
+        physical_schema = PyarrowFieldParser.from_paimon_schema(
+            builder.read_type())
+        effective_schema = _effective_masked_schema(
+            physical_schema, masking)
+        visible_paths = (
+            None if self._name_paths is None
+            else self._name_paths[:len(self._schema)]
+        )
+        self._schema = _project_effective_schema(
+            self._schema, visible_paths, effective_schema, masking)
+        self._fetch_schema = _project_effective_schema(
+            self._fetch_schema, self._name_paths,
+            effective_schema, masking)
+        self._split_ranges = [
+            self._row_ranges(split) for split in self._splits]
+        self._range_intervals = sorted(
+            (row_range.from_, row_range.to, split_index)
+            for split_index, ranges in enumerate(self._split_ranges)
+            for row_range in ranges
+        )
+        self._range_starts = [
+            interval[0] for interval in self._range_intervals]
+        self._range_max_ends = []
+        max_end = -1
+        for _, end, _ in self._range_intervals:
+            max_end = max(max_end, end)
+            self._range_max_ends.append(max_end)
+
+    @property
+    def schema(self):
+        return self._schema
+
+    @staticmethod
+    def _row_ranges(split):
+        if isinstance(split, QueryAuthSplit):
+            split = split.split
+        if isinstance(split, IndexedSplit):
+            ranges = split.row_ranges()
+        else:
+            ranges = [
+                data_file.row_id_range()
+                for data_file in split.files
+                if data_file.row_id_range() is not None
+            ]
+        return Range.sort_and_merge_overlap(ranges, True)
+
+    def fetch(self, row_ids):
+        if not row_ids:
+            return pa.Table.from_arrays(
+                [pa.array([], type=field.type) for field in self._schema],
+                schema=self._schema,
+            )
+
+        wanted = Range.sort_and_merge_overlap(
+            [Range(row_id, row_id) for row_id in set(row_ids)], True)
+        selected_splits = []
+        for split_index in self._find_splits(wanted):
+            original = self._splits[split_index]
+            auth_result = None
+            split = original
+            if isinstance(split, QueryAuthSplit):
+                auth_result = split.auth_result
+                split = split.split
+            if isinstance(split, IndexedSplit):
+                split = split.data_split()
+            allowed = Range.and_(wanted, self._split_ranges[split_index])
+            if not allowed:
+                continue
+            indexed = IndexedSplit(
+                split,
+                allowed,
+                exact_merged_row_count=sum(r.count() for r in allowed),
+            )
+            if auth_result is not None:
+                indexed = QueryAuthSplit(indexed, auth_result)
+            selected_splits.append(indexed)
+
+        arrow = self._project_fetch(self._read.to_arrow(selected_splits))
+        found = arrow[self._row_id_name].to_pylist()
+        positions = {}
+        for index, row_id in enumerate(found):
+            if row_id in positions:
+                raise RuntimeError(
+                    "Duplicate row id %r in aligned scan." % row_id)
+            positions[row_id] = index
+        missing = [row_id for row_id in row_ids if row_id not in positions]
+        if missing:
+            raise RuntimeError(
+                "Aligned row ids disappeared from pinned snapshot: %r."
+                % missing
+            )
+        take = pa.array(
+            [positions[row_id] for row_id in row_ids], type=pa.int64())
+        return arrow.select(self._schema.names).take(take)
+
+    def _find_splits(self, ranges):
+        split_indices = set()
+        for row_range in ranges:
+            right = bisect_right(self._range_starts, row_range.to)
+            left = bisect_left(
+                self._range_max_ends, row_range.from_, 0, right)
+            for position in range(left, right):
+                _, end, split_index = self._range_intervals[position]
+                if end >= row_range.from_:
+                    split_indices.add(split_index)
+        return sorted(split_indices)
+
+    def _project_fetch(self, arrow):
+        if self._name_paths is None:
+            return arrow
+        arrays = []
+        for path in self._name_paths:
+            array = arrow[path[0]]
+            for name in path[1:]:
+                index = array.type.get_field_index(name)
+                if index < 0:
+                    raise KeyError("Nested field %r does not exist." % name)
+                array = array.flatten()[index]
+            arrays.append(array)
+        return pa.Table.from_arrays(arrays, schema=self._fetch_schema)
+
+
+def _arrow_rows(table):
+    if hasattr(table, "to_pylist"):
+        return table.to_pylist()
+    columns = table.to_pydict()
+    return [
+        {name: columns[name][index] for name in table.column_names}
+        for index in range(table.num_rows)
+    ]
+
+
+def _metadata_rows(table, on, time_type):
+    rows = _arrow_rows(table)
+    if pa.types.is_timestamp(time_type):
+        times = table[on].combine_chunks()
+        for index, row in enumerate(rows):
+            row[_TIME_KEY] = times[index].value
+    else:
+        for row in rows:
+            row[_TIME_KEY] = row[on]
+    return rows
+
+
+def _validate_pinned_tag(query):
+    tag_name = getattr(query, "_temporal_tag_name", None)
+    if tag_name is None:
+        return
+    tag = query._table.tag_manager().get(tag_name)
+    if tag is None or tag.id != query._temporal_snapshot_id:
+        raise RuntimeError(
+            "Tag %r changed after temporal alignment was created." % tag_name)
+
+
+def _plan_with_internal_row_id(
+        builder, expected_snapshot_id, empty=False):
+    scan = builder.new_scan()
+    if empty:
+        auth_result = resolve_auth_result(
+            getattr(scan, "_query_auth_fn", None), scan._read_type)
+        masking = _masking_rules([auth_result])
+        if _ROW_ID in masking:
+            raise ValueError(
+                "Temporal alignment cannot use a query that masks _ROW_ID.")
+        return [], masking
+
+    auth_results = []
+    query_auth = getattr(scan, "_query_auth_fn", None)
+    if query_auth is not None:
+        def capture_auth(select):
+            result = query_auth(select)
+            auth_results.append(result)
+            return result
+
+        scan._query_auth_fn = capture_auth
+    plan = scan.plan()
+    if plan.snapshot_id != expected_snapshot_id:
+        raise RuntimeError(
+            "Temporal input changed from snapshot %r to %r during planning."
+            % (expected_snapshot_id, plan.snapshot_id)
+        )
+    splits = plan.splits()
+    auth_results.extend(
+        split.auth_result for split in splits
+        if isinstance(split, QueryAuthSplit)
+    )
+    masking = _masking_rules(auth_results)
+    if _ROW_ID in masking:
+        raise ValueError(
+            "Temporal alignment cannot use a query that masks _ROW_ID.")
+    return splits, masking
+
+
+def _masking_rules(auth_results):
+    masking = None
+    for auth_result in auth_results:
+        if auth_result is None:
+            continue
+        current = dict(
+            getattr(auth_result, "column_masking", None) or {})
+        if masking is None:
+            masking = current
+        elif current != masking:
+            raise RuntimeError(
+                "Column masking rules changed during query planning.")
+    parsed = {}
+    for name, rule in (masking or {}).items():
+        if not rule:
+            continue
+        transform = json.loads(rule)
+        if transform is not None:
+            parsed[name] = transform
+    return parsed
+
+
+def _with_active_masking(splits, targets):
+    active = set(targets)
+    result = []
+    for split in splits:
+        if not isinstance(split, QueryAuthSplit):
+            result.append(split)
+            continue
+        auth = split.auth_result
+        masking = getattr(auth, "column_masking", None) or {}
+        restricted = {
+            name: rule for name, rule in masking.items() if name in active
+        }
+        if restricted == masking:
+            result.append(split)
+            continue
+        auth = TableQueryAuthResult(
+            filter=getattr(auth, "filter", None),
+            column_masking=restricted or None,
+        )
+        result.append(
+            QueryAuthSplit(split.split, auth)
+            if auth.has_restrictions else split.split
+        )
+    return result
+
+
+def _mask_dependencies(masking, targets, table_schema):
+    dependencies = set()
+    ordered_targets = list(dict.fromkeys(targets))
+    readable = set(ordered_targets)
+    pending = list(ordered_targets)
+    while pending:
+        target = pending.pop(0)
+        transform = masking.get(target)
+        if transform is None:
+            continue
+        for name in _collect_all_field_refs_from_transform(transform):
+            if name not in table_schema.names:
+                raise ValueError(
+                    "Column masking for %r refers to unknown field %r."
+                    % (target, name)
+                )
+            if name not in readable:
+                readable.add(name)
+                dependencies.add(name)
+                pending.append(name)
+    return [
+        name for name in table_schema.names if name in dependencies
+    ]
+
+
+def _effective_masked_schema(schema, masking):
+    if not masking:
+        return schema
+    batch = pa.RecordBatch.from_arrays([
+        pa.array([], type=field.type) for field in schema
+    ], schema=schema)
+    fields = []
+    for field in schema:
+        transform = masking.get(field.name)
+        if transform is None:
+            fields.append(field)
+            continue
+        masked = _apply_predicate_transform(
+            transform, batch, null_type=field.type)
+        fields.append(pa.field(
+            field.name, masked.type, nullable=True,
+            metadata=field.metadata))
+    return pa.schema(fields, metadata=schema.metadata)
+
+
+def _project_effective_schema(
+        schema, name_paths, effective_schema, masking):
+    if not masking:
+        return schema
+    paths = name_paths or [(field.name,) for field in schema]
+    fields = []
+    for field, path in zip(schema, paths):
+        target = path[0]
+        if not masking.get(target):
+            fields.append(field)
+            continue
+        masked_field = effective_schema.field(target)
+        masked_type = masked_field.type
+        for name in path[1:]:
+            if not pa.types.is_struct(masked_type):
+                raise TypeError(
+                    "Column masking for %r no longer produces the struct "
+                    "required by nested projection %r."
+                    % (target, ".".join(path))
+                )
+            index = masked_type.get_field_index(name)
+            if index < 0:
+                raise TypeError(
+                    "Column masking for %r does not produce nested field %r."
+                    % (target, name)
+                )
+            masked_type = masked_type[index].type
+        fields.append(pa.field(
+            field.name, masked_type, nullable=True,
+            metadata=field.metadata))
+    return pa.schema(fields, metadata=schema.metadata)
+
+
+def _match_index(times, target, method, start=0, end=None):
+    end = len(times) if end is None else end
+    if start >= end:
+        return None
+    position = bisect_left(times, target, start, end)
+    if method == "backward":
+        position = bisect_right(times, target, start, end)
+        return position - 1 if position > start else None
+    if method == "forward":
+        return position if position < end else None
+    if method == "nearest":
+        if position < end and times[position] == target:
+            return bisect_right(times, target, position, end) - 1
+        if position == start:
+            return start
+        if position == end:
+            return end - 1
+        before = position - 1
+        before_value = _python_scalar(times[before])
+        after_value = _python_scalar(times[position])
+        if target - before_value <= after_value - target:
+            return before
+        return position
+    raise ValueError("Unknown temporal match method %r." % method)
+
+
+def _time_search_keys(values, data_type):
+    if pa.types.is_timestamp(data_type):
+        values = pc.cast(values, pa.int64())
+    return values.to_numpy(zero_copy_only=False)
+
+
+def _time_tolerance_key(tolerance, data_type):
+    if tolerance is None or not pa.types.is_timestamp(data_type):
+        return tolerance
+    return pa.scalar(tolerance, type=pa.duration(data_type.unit)).value
+
+
+def _python_scalar(value):
+    item = getattr(value, "item", None)
+    return item() if item is not None else value
+
+
+def _validate_tolerance(tolerance, data_type):
+    if tolerance is None:
+        return
+    if pa.types.is_timestamp(data_type):
+        if not isinstance(tolerance, timedelta):
+            raise TypeError(
+                "Timestamp alignment tolerance must be datetime.timedelta.")
+        return
+    if isinstance(tolerance, timedelta):
+        raise TypeError("Numeric alignment tolerance must be numeric.")
+    if pa.types.is_integer(data_type) and tolerance > _MAX_INT64:
+        raise ValueError(
+            "Integer alignment tolerance cannot exceed int64 maximum.")
+
+
+def _delta_type(data_type):
+    if pa.types.is_timestamp(data_type):
+        return pa.duration(data_type.unit)
+    if pa.types.is_integer(data_type):
+        return pa.int64()
+    if pa.types.is_floating(data_type):
+        return pa.float64()
+    raise TypeError(
+        "Temporal columns must be integer, floating point, or timestamp; "
+        "got %s."
+        % data_type
+    )
+
+
+def _validate_group_type(name, data_type):
+    if pa.types.is_nested(data_type) or pa.types.is_null(data_type):
+        raise TypeError(
+            "Grouping column %r must have a scalar type; got %s."
+            % (name, data_type)
+        )
+
+
+def _require_columns(schema, columns, label):
+    missing = [name for name in columns if name not in schema.names]
+    if missing:
+        raise ValueError(
+            "%s is missing temporal columns %r." % (label, missing))
diff --git a/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py 
b/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py
index c84a1383ad..0d067560b4 100644
--- a/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py
+++ b/paimon-python/pypaimon/read/reader/format_pyarrow_reader.py
@@ -80,7 +80,7 @@ class _FileFormatDatasetCache:
         self._loads = {}
         self._lock = threading.Lock()
 
-    def get_or_load(self, key: Tuple[Any, str, str], loader: Callable[[], Any],
+    def get_or_load(self, key: Tuple[Any, ...], loader: Callable[[], Any],
                     size_estimator: Callable[[Any], Optional[int]]):
         with self._lock:
             entry = self._entries.get(key)
@@ -109,7 +109,8 @@ class _FileFormatDatasetCache:
             raise
 
         with self._lock:
-            if estimated_size is not None:
+            if (estimated_size is not None
+                    and estimated_size <= self.max_size):
                 estimated_size = max(1, estimated_size)
                 self._entries[key] = (dataset, estimated_size)
                 self.estimated_size += estimated_size
@@ -120,6 +121,23 @@ class _FileFormatDatasetCache:
             self._loads.pop(key, None)
         return dataset
 
+    def get(self, key):
+        with self._lock:
+            entry = self._entries.get(key)
+            if entry is None:
+                return None
+            self._entries.move_to_end(key)
+            return entry[0]
+
+    def put(self, key, dataset, estimated_size):
+        estimated_size = max(1, estimated_size)
+        with self._lock:
+            if estimated_size > self.max_size or key in self._entries:
+                return
+            self._entries[key] = (dataset, estimated_size)
+            self.estimated_size += estimated_size
+            self._evict()
+
     def resize(self, max_size: int):
         with self._lock:
             self.max_size = max_size
@@ -133,6 +151,32 @@ class _FileFormatDatasetCache:
             self.estimated_size -= evicted_size
 
 
+class _DecodedRowGroupCache:
+    def __init__(self, max_size: int):
+        self._cache = _FileFormatDatasetCache(max_size)
+
+    def iter_or_load(self, key, loader):
+        cached = self._cache.get(key)
+        if cached is not None:
+            yield from cached
+            return
+
+        batches = []
+        size = 0
+        for batch in loader():
+            # Yield while loading so one oversized row group is never
+            # materialized in full merely to discover that it cannot fit.
+            if batches is not None:
+                size += batch.nbytes
+                if size <= self._cache.max_size:
+                    batches.append(batch)
+                else:
+                    batches = None
+            yield batch
+        if batches is not None:
+            self._cache.put(key, batches, size)
+
+
 _FILE_FORMAT_DATASET_CACHE = None
 _FILE_FORMAT_DATASET_CACHE_LOCK = threading.Lock()
 _FILE_FORMAT_DATASET_CACHE_PID = os.getpid()
@@ -253,9 +297,14 @@ class FormatPyArrowReader(RecordBatchReader):
                  nested_name_paths: Optional[List[List[str]]] = None,
                  predicate_field_names: Optional[Set[str]] = None,
                  row_indices: Optional[List[int]] = None,
-                 row_ranges: Optional[List[Tuple[int, int]]] = None):
+                 row_ranges: Optional[List[Tuple[int, int]]] = None,
+                 row_group_cache: Optional[_DecodedRowGroupCache] = None):
         self._predicate_field_names = predicate_field_names or set()
         file_path_for_pyarrow = file_io.to_filesystem_path(file_path)
+        self._row_group_cache = row_group_cache
+        self._row_group_cache_filesystem = _FilesystemIdentity(
+            file_io.filesystem)
+        self._row_group_cache_path = file_path_for_pyarrow
         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)
@@ -397,18 +446,32 @@ class FormatPyArrowReader(RecordBatchReader):
     def _iter_row_group_batches(self):
         columns = self._row_group_read_columns()
         for row_group in self._surviving_row_group_ids():
-            for batch in self._parquet_file.iter_batches(
-                    row_groups=[row_group],
-                    columns=columns,
-                    batch_size=self._scan_batch_size):
+            if (self._row_group_cache is not None
+                    and self._selected_parquet_row_groups is not None):
+                key = (
+                    self._row_group_cache_filesystem,
+                    self._row_group_cache_path,
+                    row_group,
+                    tuple(columns),
+                    self._scan_batch_size,
+                )
+                batches = self._row_group_cache.iter_or_load(
+                    key,
+                    lambda: self._read_parquet_row_group_batches(
+                        row_group, columns),
+                )
+            else:
+                batches = self._read_parquet_row_group_batches(
+                    row_group, columns)
+            for batch in batches:
                 if self._has_nested_path:
-                    batches = [batch]
+                    filtered_batches = [batch]
                     if self._scan_filter is not None:
                         table = ds.dataset(
                             pa.Table.from_batches([batch])
                         ).scanner(filter=self._scan_filter).to_table()
-                        batches = table.to_batches()
-                    for filtered in batches:
+                        filtered_batches = table.to_batches()
+                    for filtered in filtered_batches:
                         out = self._select_nested_fields(filtered)
                         if out.num_rows:
                             yield out
@@ -425,6 +488,13 @@ class FormatPyArrowReader(RecordBatchReader):
                     if out.num_rows:
                         yield out
 
+    def _read_parquet_row_group_batches(self, row_group, columns):
+        return self._parquet_file.iter_batches(
+            row_groups=[row_group],
+            columns=columns,
+            batch_size=self._scan_batch_size,
+        )
+
     def _row_group_read_columns(self):
         if self._has_nested_path:
             existing = set(self.existing_fields)
diff --git a/paimon-python/pypaimon/read/split_read.py 
b/paimon-python/pypaimon/read/split_read.py
index f00976f1f0..c5300a620f 100644
--- a/paimon-python/pypaimon/read/split_read.py
+++ b/paimon-python/pypaimon/read/split_read.py
@@ -152,6 +152,7 @@ class SplitRead(ABC):
         self.nested_name_paths = nested_name_paths
         self.limit = limit
         self._blob_parallelism = 1
+        self._parquet_row_group_cache = None
         # Snapshot the raw value-side schema before _create_key_value_fields
         # wraps it, so MergeFileSplitRead can hand per-value-field nullable
         # flags to merge functions that enforce NOT-NULL on every add().
@@ -393,7 +394,8 @@ class SplitRead(ABC):
                 options=self.table.options,
                 nested_name_paths=ordered_nested_paths,
                 predicate_field_names=predicate_fields,
-                row_ranges=parquet_row_ranges)
+                row_ranges=parquet_row_ranges,
+                row_group_cache=self._parquet_row_group_cache)
         elif file_format == CoreOptions.FILE_FORMAT_ROW:
             if has_nested:
                 raise NotImplementedError(
diff --git a/paimon-python/pypaimon/read/table_read.py 
b/paimon-python/pypaimon/read/table_read.py
index 51465486ed..980d8fdd9c 100644
--- a/paimon-python/pypaimon/read/table_read.py
+++ b/paimon-python/pypaimon/read/table_read.py
@@ -160,6 +160,7 @@ class TableRead:
         self.nested_name_paths = nested_name_paths
         self.limit = limit
         self._read_parallelism = self.table.options.read_parallelism()
+        self._parquet_row_group_cache = None
 
     def to_iterator(self, splits: List[Split]) -> Iterator:
         limit = self.limit
@@ -851,6 +852,7 @@ class TableRead:
             post_filter_after_inline,
         )
         sr._blob_parallelism = blob_parallelism
+        sr._parquet_row_group_cache = self._parquet_row_group_cache
         return sr
 
     def _build_split_read(self, split: Split, read_type=None,
diff --git a/paimon-python/pypaimon/tests/multimodal_temporal_test.py 
b/paimon-python/pypaimon/tests/multimodal_temporal_test.py
new file mode 100644
index 0000000000..2cb0fe7c40
--- /dev/null
+++ b/paimon-python/pypaimon/tests/multimodal_temporal_test.py
@@ -0,0 +1,1411 @@
+# 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 json
+import os
+import shutil
+import tempfile
+import unittest
+from datetime import datetime, timedelta
+from unittest import mock
+
+import pyarrow as pa
+import pypaimon.multimodal as pmm
+from pypaimon.multimodal import temporal
+from pypaimon.catalog.table_query_auth import TableQueryAuthResult
+from pypaimon.read.reader.format_pyarrow_reader import FormatPyArrowReader
+from pypaimon.read.scanner.file_scanner import FileScanner
+
+
+class MultimodalTemporalTest(unittest.TestCase):
+
+    def setUp(self):
+        self.temp_dir = tempfile.mkdtemp(prefix="pypaimon_temporal_")
+        self.conn = pmm.connect(options={
+            "warehouse": os.path.join(self.temp_dir, "warehouse"),
+        })
+
+    def tearDown(self):
+        shutil.rmtree(self.temp_dir, ignore_errors=True)
+
+    def test_alignment_requires_an_explicit_group_boundary(self):
+        table = self._table("missing_group", {
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        with self.assertRaisesRegex(ValueError, "grouping column"):
+            pmm.join_asof(
+                table.scan(),
+                table.scan(),
+                on="event_time",
+                by=(),
+                direction="nearest",
+                tolerance=0,
+            )
+
+    def test_alignment_preserves_payload_names(self):
+        anchors = self._table("audit_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        samples = self._table("audit_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "valid": pa.bool_(),
+            "matched_time": pa.int64(),
+            "time_delta": pa.int64(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        samples.add([{
+            "episode_id": 1,
+            "event_time": 100,
+            "valid": False,
+            "matched_time": 7,
+            "time_delta": 8,
+        }])
+
+        row = pmm.join_asof(
+            anchors.scan(),
+            samples.scan().select([
+                "valid", "matched_time", "time_delta"
+            ]),
+            on="event_time",
+            by="episode_id",
+            direction="nearest",
+            tolerance=0,
+        ).to_list()[0]
+
+        self.assertFalse(row["valid"])
+        self.assertEqual(7, row["matched_time"])
+        self.assertEqual(8, row["time_delta"])
+
+    def test_alignment_handles_duplicate_right_timestamps(self):
+        anchors = self._table("duplicate_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        samples = self._table("duplicate_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        samples.add([
+            {"episode_id": 1, "event_time": 100, "value": 1},
+            {"episode_id": 1, "event_time": 100, "value": 2},
+        ])
+
+        for direction, expected in (
+                ("backward", 2), ("forward", 1), ("nearest", 2)):
+            with self.subTest(direction=direction):
+                row = pmm.join_asof(
+                    anchors.scan(), samples.scan().select("value"),
+                    on="event_time", by="episode_id",
+                    direction=direction, tolerance=0,
+                ).to_list()[0]
+                self.assertEqual(expected, row["value"])
+
+    def test_nearest_uses_candidate_side_for_duplicate_timestamps(self):
+        anchors = self._table("duplicate_nearest_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        samples = self._table("duplicate_nearest_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([
+            {"episode_id": 1, "event_time": 90},
+            {"episode_id": 1, "event_time": 110},
+        ])
+        samples.add([
+            {"episode_id": 1, "event_time": 100, "value": 1},
+            {"episode_id": 1, "event_time": 100, "value": 2},
+        ])
+
+        rows = pmm.join_asof(
+            anchors.scan(), samples.scan().select("value"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=20,
+        ).to_list()
+
+        self.assertEqual(
+            {90: 1, 110: 2},
+            {row["event_time"]: row["value"] for row in rows},
+        )
+
+    def test_alignment_can_return_matched_timestamp(self):
+        anchors = self._table("timestamp_output_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        samples = self._table("timestamp_output_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        samples.add([
+            {"episode_id": 1, "event_time": 95, "value": 7},
+        ])
+
+        row = pmm.join_asof(
+            anchors.scan(),
+            samples.scan().select(["event_time", "value"]),
+            on="event_time", by="episode_id",
+            direction="backward", suffix="_matched",
+        ).to_list()[0]
+
+        self.assertEqual(95, row["event_time_matched"])
+        self.assertEqual(5, row["event_time"] - row["event_time_matched"])
+
+    def test_alignment_rejects_nested_group_keys(self):
+        group_type = pa.struct([pa.field("part", pa.int32())])
+        anchors = self._table("nested_group_anchors", {
+            "group": group_type,
+            "event_time": pa.int64(),
+        })
+        samples = self._table("nested_group_samples", {
+            "group": group_type,
+            "event_time": pa.int64(),
+        })
+
+        with self.assertRaisesRegex(TypeError, "must have a scalar type"):
+            pmm.join_asof(
+                anchors.scan(),
+                samples.scan(),
+                on="event_time",
+                by="group",
+                direction="nearest",
+                tolerance=0,
+            )
+
+    def test_aligns_named_sources_in_episode_local_batches(self):
+        actions = self._table("actions", {
+            "episode_id": pa.string(),
+            "event_time": pa.int64(),
+            "action": pa.int32(),
+        })
+        images = self._table("images", {
+            "episode_id": pa.string(),
+            "event_time": pa.int64(),
+            "camera": pa.string(),
+            "image": pa.string(),
+        })
+        states = self._table("states", {
+            "episode_id": pa.string(),
+            "event_time": pa.int64(),
+            "state": pa.int32(),
+        })
+        commands = self._table("commands", {
+            "episode_id": pa.string(),
+            "event_time": pa.int64(),
+            "command": pa.string(),
+        })
+        actions.add([
+            {"episode_id": "ep-2", "event_time": 100, "action": 4},
+            {"episode_id": "ep-1", "event_time": 300, "action": 3},
+            {"episode_id": "ep-1", "event_time": 100, "action": 1},
+            {"episode_id": "ep-1", "event_time": 200, "action": 2},
+        ])
+        images.add([
+            {"episode_id": "ep-1", "event_time": 90,
+             "camera": "left", "image": "early"},
+            {"episode_id": "ep-1", "event_time": 90,
+             "camera": "right", "image": "ignored"},
+            {"episode_id": "ep-1", "event_time": 110,
+             "camera": "left", "image": "late"},
+            {"episode_id": "ep-1", "event_time": 215,
+             "camera": "left", "image": "middle"},
+            {"episode_id": "ep-2", "event_time": 99,
+             "camera": "left", "image": "other"},
+        ])
+        states.add([
+            {"episode_id": "ep-1", "event_time": 80, "state": 8},
+            {"episode_id": "ep-1", "event_time": 190, "state": 19},
+            {"episode_id": "ep-2", "event_time": 95, "state": 95},
+        ])
+        commands.add([
+            {"episode_id": "ep-1", "event_time": 100, "command": "open"},
+            {"episode_id": "ep-1", "event_time": 220, "command": "close"},
+            {"episode_id": "ep-2", "event_time": 100, "command": "hold"},
+        ])
+
+        aligned = pmm.join_asof(
+            actions.scan().select(["episode_id", "event_time", "action"]),
+            images.scan().where("camera = 'left'").select("image"),
+            on="event_time",
+            by="episode_id",
+            direction="nearest",
+            tolerance=20,
+        ).join_asof(
+            states.scan().select("state"),
+            direction="backward",
+            tolerance=25,
+        ).join_asof(
+            commands.scan().select("command"),
+            direction="nearest",
+            tolerance=0,
+        ).join_asof(
+            commands.scan().select("command"),
+            direction="forward",
+            tolerance=25,
+            suffix="_next",
+        )
+        reader = aligned.to_arrow_batch_reader(batch_size=2)
+        batches = list(reader)
+        rows = pa.Table.from_batches(batches).to_pylist()
+
+        self.assertEqual([2, 2], [batch.num_rows for batch in batches])
+        rows.sort(key=lambda row: (row["episode_id"], row["event_time"]))
+        self.assertEqual(
+            [("ep-1", 100), ("ep-1", 200), ("ep-1", 300), ("ep-2", 100)],
+            [(row["episode_id"], row["event_time"]) for row in rows],
+        )
+        # Equal-distance nearest ties select the earlier row.
+        self.assertEqual(
+            ["early", "middle", None, "other"],
+            [row["image"] for row in rows],
+        )
+        self.assertEqual([8, 19, None, 95], [
+            row["state"] for row in rows
+        ])
+        self.assertEqual(["open", None, None, "hold"], [
+            row["command"] for row in rows
+        ])
+        self.assertEqual(["open", "close", None, "hold"], [
+            row["command_next"] for row in rows
+        ])
+
+    def test_alignment_pins_each_scan_snapshot(self):
+        anchors = self._table("pinned_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.string(),
+        })
+        secondary = self._table("pinned_secondary", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.string(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100, "value": "old"}])
+        secondary.add([
+            {"episode_id": 1, "event_time": 90, "value": "old-match"}
+        ])
+        aligned = pmm.join_asof(
+            anchors.scan(),
+            secondary.scan(),
+            on="event_time",
+            by="episode_id",
+            direction="nearest",
+            tolerance=20,
+            suffix="_secondary",
+        )
+        snapshots = aligned.resolved_snapshots
+
+        anchors.add([{"episode_id": 1, "event_time": 200, "value": "new"}])
+        secondary.add([
+            {"episode_id": 1, "event_time": 100, "value": "new-match"}
+        ])
+
+        self.assertEqual([{
+            "episode_id": 1,
+            "event_time": 100,
+            "value": "old",
+            "value_secondary": "old-match",
+        }], aligned.to_list())
+        self.assertEqual(1, snapshots["left"]["snapshot_id"])
+        self.assertEqual(1, snapshots["right_1"]["snapshot_id"])
+        self.assertTrue(
+            snapshots["left"]["table"].endswith("pinned_anchors"))
+        self.assertTrue(
+            snapshots["right_1"]["table"].endswith("pinned_secondary"))
+
+    def test_alignment_reads_tag_after_snapshot_file_expires(self):
+        anchors = self._table("tagged_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.string(),
+        })
+        secondary = self._table("tagged_secondary", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.string(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100, "value": "old"}])
+        secondary.add([
+            {"episode_id": 1, "event_time": 100, "value": "old-match"}
+        ])
+        anchors.raw_table.create_tag("v1")
+        secondary.raw_table.create_tag("v1")
+        aligned = pmm.join_asof(
+            anchors.scan(tag_name="v1"),
+            secondary.scan(tag_name="v1"),
+            on="event_time",
+            by="episode_id",
+            direction="nearest",
+            tolerance=0,
+        )
+
+        anchors.add([{"episode_id": 1, "event_time": 200, "value": "new"}])
+        secondary.add([
+            {"episode_id": 1, "event_time": 200, "value": "new-match"}
+        ])
+        for table in (anchors.raw_table, secondary.raw_table):
+            manager = table.snapshot_manager()
+            table.file_io.delete(manager.get_snapshot_path(1))
+
+        self.assertEqual([100], [
+            row["event_time"] for row in aligned.to_list()
+        ])
+        self.assertEqual(
+            "v1", aligned.resolved_snapshots["left"]["tag_name"])
+
+    def test_alignment_rejects_a_changed_tag(self):
+        table = self._table("changed_tag", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        table.add([{"episode_id": 1, "event_time": 100}])
+        table.raw_table.create_tag("v1")
+        aligned = pmm.join_asof(
+            table.scan(tag_name="v1"), table.scan(tag_name="v1"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        )
+        table.add([{"episode_id": 1, "event_time": 200}])
+        table.raw_table.replace_tag("v1")
+
+        with self.assertRaisesRegex(RuntimeError, "Tag 'v1' changed"):
+            aligned.to_list()
+
+    def test_alignment_rejects_tag_replacement_during_planning(self):
+        anchors = self._table("raced_tag_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        source = self._table("raced_tag_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        source.add([{
+            "episode_id": 1, "event_time": 100, "value": 7,
+        }])
+        source.raw_table.create_tag("v1")
+        source.add([{
+            "episode_id": 1, "event_time": 200, "value": 99,
+        }])
+        aligned = pmm.join_asof(
+            anchors.scan(), source.scan(tag_name="v1").select("value"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        )
+        original = temporal._validate_pinned_tag
+        replaced = []
+
+        def replace_after_validation(query):
+            original(query)
+            if query._temporal_tag_name == "v1" and not replaced:
+                source.raw_table.replace_tag("v1", snapshot_id=2)
+                replaced.append(True)
+
+        with mock.patch.object(
+                temporal, "_validate_pinned_tag",
+                side_effect=replace_after_validation):
+            with self.assertRaisesRegex(
+                    RuntimeError, "changed from snapshot 1 to 2"):
+                aligned.to_list()
+
+    def test_alignment_normalizes_scan_mode_when_pinning(self):
+        anchors = self.conn.create_table(
+            "latest_full_anchors",
+            schema=pa.schema([
+                pa.field("episode_id", pa.int32()),
+                pa.field("event_time", pa.int64()),
+            ]),
+            options={"scan.mode": "latest-full"},
+        )
+        secondary = self._table("latest_full_secondary", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        secondary.add([
+            {"episode_id": 1, "event_time": 100, "value": 7}
+        ])
+
+        row = pmm.join_asof(
+            anchors.scan(),
+            secondary.scan(),
+            on="event_time",
+            by="episode_id",
+            direction="nearest",
+            tolerance=0,
+        ).to_list()[0]
+
+        self.assertEqual(7, row["value"])
+
+    def test_alignment_rejects_non_finite_temporal_values(self):
+        for value in (float("nan"), float("inf"), float("-inf")):
+            with self.subTest(value=value):
+                anchors = self._table("float_anchor_%s" % id(value), {
+                    "episode_id": pa.int32(),
+                    "event_time": pa.float64(),
+                })
+                secondary = self._table("float_source_%s" % id(value), {
+                    "episode_id": pa.int32(),
+                    "event_time": pa.float64(),
+                    "value": pa.int32(),
+                })
+                anchors.add([{"episode_id": 1, "event_time": 100.0}])
+                secondary.add([
+                    {"episode_id": 1, "event_time": value, "value": 7}
+                ])
+                aligned = pmm.join_asof(
+                    anchors.scan(),
+                    secondary.scan(),
+                    on="event_time",
+                    by="episode_id",
+                    direction="nearest",
+                    tolerance=1.0,
+                )
+                with self.assertRaisesRegex(ValueError, "must be finite"):
+                    aligned.to_list()
+
+    def test_alignment_validates_tolerance_type_and_value(self):
+        table = self._table("tolerance", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        for tolerance in (float("nan"), float("inf"), -1):
+            with self.subTest(tolerance=tolerance):
+                with self.assertRaises((TypeError, ValueError)):
+                    pmm.join_asof(
+                        table.scan(), table.scan(),
+                        on="event_time", by="episode_id",
+                        direction="nearest", tolerance=tolerance,
+                    )
+        with self.assertRaisesRegex(TypeError, "Numeric alignment"):
+            pmm.join_asof(
+                table.scan(), table.scan(),
+                on="event_time", by="episode_id",
+                direction="nearest",
+                tolerance=timedelta(milliseconds=1),
+            )
+        with self.assertRaisesRegex(ValueError, "int64 maximum"):
+            pmm.join_asof(
+                table.scan(), table.scan(),
+                on="event_time", by="episode_id",
+                direction="nearest", tolerance=1 << 63,
+            )
+        with self.assertRaisesRegex(ValueError, "direction"):
+            pmm.join_asof(
+                table.scan(), table.scan(),
+                on="event_time", by="episode_id", direction="exact",
+            )
+
+    def test_alignment_keeps_internal_row_id_out_of_query_auth(self):
+        anchors = self._table("masked_row_id_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        source = self._table("masked_row_id_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        source.add([{"episode_id": 1, "event_time": 100, "value": 7}])
+        selected = []
+
+        def query_auth(select):
+            selected.append(select)
+            self.assertTrue(select is None or "_ROW_ID" not in select)
+            return None
+
+        for table in (anchors.raw_table, source.raw_table):
+            table.catalog_environment.table_query_auth = (
+                lambda options, identifier: query_auth)
+
+        rows = pmm.join_asof(
+            anchors.scan(), source.scan(),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        ).to_list()
+
+        self.assertEqual(7, rows[0]["value"])
+        self.assertTrue(selected)
+
+    def test_alignment_rejects_masked_internal_row_ids(self):
+        anchors = self._table("masked_internal_id_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        source = self._table("masked_internal_id_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        source.add([{"episode_id": 1, "event_time": 100, "value": 7}])
+        auth = TableQueryAuthResult(
+            filter=None,
+            column_masking={"_ROW_ID": json.dumps({"name": "NULL"})},
+        )
+        anchors.raw_table.catalog_environment.table_query_auth = (
+            lambda options, identifier: lambda select: auth)
+
+        aligned = pmm.join_asof(
+            anchors.scan(), source.scan(),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        )
+        with self.assertRaisesRegex(ValueError, "masks _ROW_ID"):
+            aligned.to_list()
+
+    def test_alignment_preserves_masked_output_schema(self):
+        anchors = self.conn.create_table(
+            "masked_schema_anchors",
+            schema=pa.schema([
+                pa.field("episode_id", pa.int32()),
+                pa.field("event_time", pa.int64()),
+                pa.field("secret", pa.int32(), nullable=False),
+            ]),
+        )
+        source = self._table("masked_schema_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "sample": pa.int32(),
+        })
+        anchors.add([{
+            "episode_id": 1, "event_time": 100, "secret": 7,
+        }])
+        source.add([{
+            "episode_id": 1, "event_time": 100, "sample": 1,
+        }])
+        auth = TableQueryAuthResult(
+            filter=None,
+            column_masking={"secret": json.dumps({"name": "NULL"})},
+        )
+        anchors.raw_table.catalog_environment.table_query_auth = (
+            lambda options, identifier: lambda select: auth)
+
+        reader = pmm.join_asof(
+            anchors.scan(), source.scan().select("sample"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        ).to_arrow_batch_reader()
+        result = reader.read_all()
+
+        self.assertTrue(result.schema.field("secret").nullable)
+        self.assertIsNone(result["secret"][0].as_py())
+
+    def test_alignment_preserves_type_changing_masked_output(self):
+        anchors = self._table("cast_mask_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "anchor_value": pa.int32(),
+        })
+        source = self._table("cast_mask_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "source_value": pa.int32(),
+        })
+        anchors.add([{
+            "episode_id": 1, "event_time": 100, "anchor_value": 7,
+        }])
+        source.add([{
+            "episode_id": 1, "event_time": 100, "source_value": 8,
+        }])
+
+        for table, name in (
+                (anchors.raw_table, "anchor_value"),
+                (source.raw_table, "source_value")):
+            auth = TableQueryAuthResult(
+                filter=None,
+                column_masking={name: json.dumps({
+                    "name": "CAST",
+                    "fieldRef": {"index": 2, "name": name, "type": "INT"},
+                    "type": "STRING",
+                })},
+            )
+            table.catalog_environment.table_query_auth = (
+                lambda options, identifier, result=auth:
+                lambda select: result)
+
+        result = pmm.join_asof(
+            anchors.scan(), source.scan().select("source_value"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        ).to_arrow()
+
+        self.assertEqual(pa.string(), result.schema.field("anchor_value").type)
+        self.assertEqual(pa.string(), result.schema.field("source_value").type)
+        self.assertEqual("7", result["anchor_value"][0].as_py())
+        self.assertEqual("8", result["source_value"][0].as_py())
+
+    def test_alignment_preserves_masked_schema_for_empty_source(self):
+        anchors = self._table("empty_mask_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        source = self._table("empty_mask_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        auth = TableQueryAuthResult(
+            filter=None,
+            column_masking={"value": json.dumps({
+                "name": "CAST",
+                "fieldRef": {
+                    "index": 2, "name": "value", "type": "INT",
+                },
+                "type": "STRING",
+            })},
+        )
+        source.raw_table.catalog_environment.table_query_auth = (
+            lambda options, identifier: lambda select: auth)
+
+        result = pmm.join_asof(
+            anchors.scan(), source.scan().select("value"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        ).to_arrow()
+
+        self.assertEqual(pa.string(), result.schema.field("value").type)
+        self.assertIsNone(result["value"][0].as_py())
+
+    def test_alignment_rejects_type_changing_key_masks(self):
+        anchors = self._table("cast_key_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        source = self._table("cast_key_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        source.add([{
+            "episode_id": 1, "event_time": 100, "value": 7,
+        }])
+
+        for name, index, field_type in (
+                ("episode_id", 0, "INT"),
+                ("event_time", 1, "BIGINT")):
+            with self.subTest(name=name):
+                auth = TableQueryAuthResult(
+                    filter=None,
+                    column_masking={name: json.dumps({
+                        "name": "CAST",
+                        "fieldRef": {
+                            "index": index, "name": name, "type": field_type,
+                        },
+                        "type": "STRING",
+                    })},
+                )
+                source.raw_table.catalog_environment.table_query_auth = (
+                    lambda options, identifier, result=auth:
+                    lambda select: result)
+
+                with self.assertRaisesRegex(TypeError, name):
+                    pmm.join_asof(
+                        anchors.scan(), source.scan().select("value"),
+                        on="event_time", by="episode_id",
+                        direction="nearest", tolerance=0,
+                    ).to_list()
+
+    def test_alignment_supports_cross_column_key_masks(self):
+        anchors = self._table("cross_mask_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        source = self._table("cross_mask_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "allowed_episode_id": pa.int32(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        source.add([{
+            "episode_id": 99,
+            "event_time": 100,
+            "allowed_episode_id": 1,
+            "value": 7,
+        }])
+        auth = TableQueryAuthResult(
+            filter=None,
+            column_masking={"episode_id": json.dumps({
+                "name": "FIELD_REF",
+                "fieldRef": {
+                    "index": 2,
+                    "name": "allowed_episode_id",
+                    "type": "INT",
+                },
+            })},
+        )
+        selected = []
+
+        def query_auth(select):
+            selected.append(select)
+            return auth
+
+        source.raw_table.catalog_environment.table_query_auth = (
+            lambda options, identifier: query_auth)
+
+        rows = pmm.join_asof(
+            anchors.scan(), source.scan().select("value"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        ).to_list()
+
+        self.assertEqual(7, rows[0]["value"])
+        self.assertTrue(selected)
+        self.assertTrue(all(
+            select is None or "allowed_episode_id" not in select
+            for select in selected
+        ))
+
+    def test_alignment_does_not_mask_internal_key_dependencies(self):
+        anchors = self._table("dependency_mask_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "allowed_episode_id": pa.int32(),
+        })
+        source = self._table("dependency_mask_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{
+            "episode_id": 99,
+            "event_time": 100,
+            "allowed_episode_id": 1,
+        }])
+        source.add([{
+            "episode_id": 1, "event_time": 100, "value": 7,
+        }])
+        auth = TableQueryAuthResult(
+            filter=None,
+            column_masking={
+                "episode_id": json.dumps({
+                    "name": "FIELD_REF",
+                    "fieldRef": {
+                        "index": 2,
+                        "name": "allowed_episode_id",
+                        "type": "INT",
+                    },
+                }),
+                "allowed_episode_id": json.dumps({
+                    "name": "CAST",
+                    "fieldRef": {
+                        "index": 2,
+                        "name": "allowed_episode_id",
+                        "type": "INT",
+                    },
+                    "type": "STRING",
+                }),
+            },
+        )
+        anchors.raw_table.catalog_environment.table_query_auth = (
+            lambda options, identifier: lambda select: auth)
+
+        result = pmm.join_asof(
+            anchors.scan(), source.scan().select("value"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        ).to_arrow()
+
+        self.assertEqual(7, result["value"][0].as_py())
+        self.assertEqual(1, result["episode_id"][0].as_py())
+        self.assertEqual("1", result["allowed_episode_id"][0].as_py())
+        self.assertEqual(
+            pa.string(), result.schema.field("allowed_episode_id").type)
+
+    def test_alignment_matches_masking_reader_rule_semantics(self):
+        anchors = self._table("mask_semantics_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        source = self._table("mask_semantics_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "first": pa.string(),
+            "second": pa.string(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        source.add([{
+            "episode_id": 1,
+            "event_time": 100,
+            "first": "a",
+            "second": "b",
+        }])
+        auth = [TableQueryAuthResult(
+            filter=None,
+            column_masking={
+                "first": json.dumps({
+                    "name": "FIELD_REF",
+                    "fieldRef": {
+                        "index": 3, "name": "second", "type": "STRING",
+                    },
+                }),
+                "second": json.dumps({
+                    "name": "FIELD_REF",
+                    "fieldRef": {
+                        "index": 2, "name": "first", "type": "STRING",
+                    },
+                }),
+            },
+        )]
+        source.raw_table.catalog_environment.table_query_auth = (
+            lambda options, identifier: lambda select: auth[0])
+
+        row = pmm.join_asof(
+            anchors.scan(), source.scan().select(["first", "second"]),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        ).to_list()[0]
+        self.assertEqual(("b", "a"), (row["first"], row["second"]))
+
+        auth[0] = TableQueryAuthResult(
+            filter=None, column_masking={"first": "null"})
+        row = pmm.join_asof(
+            anchors.scan(), source.scan().select("first"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        ).to_list()[0]
+        self.assertEqual("a", row["first"])
+
+    def test_alignment_rejects_incremental_scans(self):
+        anchors = self.conn.create_table(
+            "incremental_anchors",
+            schema=pa.schema([
+                pa.field("episode_id", pa.int32()),
+                pa.field("event_time", pa.int64()),
+            ]),
+            options={
+                "scan.mode": "incremental",
+                "incremental-between-timestamp": "0,9999999999999",
+            },
+        )
+        source = self._table("incremental_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        source.add([{
+            "episode_id": 1, "event_time": 100, "value": 7,
+        }])
+
+        with self.assertRaisesRegex(
+                ValueError, "join_asof.*incremental"):
+            pmm.join_asof(
+                anchors.scan(), source.scan().select("value"),
+                on="event_time", by="episode_id",
+                direction="nearest", tolerance=0,
+            ).to_list()
+
+    def test_alignment_supports_zero_output_columns(self):
+        anchors = self._table("zero_output_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        source = self._table("zero_output_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        source.add([{"episode_id": 1, "event_time": 100}])
+
+        result = pmm.join_asof(
+            anchors.scan().select("missing"),
+            source.scan().select("episode_id"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        )
+
+        self.assertEqual([], result.schema.names)
+        self.assertEqual([{}], result.to_list())
+
+    def test_alignment_supports_nested_projections(self):
+        anchors = self._table("nested_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "metadata": pa.struct([pa.field("value", pa.int32())]),
+        })
+        secondary = self._table("nested_secondary", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "payload": pa.struct([pa.field("value", pa.int32())]),
+        })
+        anchors.add([
+            {"episode_id": 1, "event_time": 100,
+             "metadata": {"value": 1}},
+            {"episode_id": 1, "event_time": 200, "metadata": None},
+        ])
+        secondary.add([
+            {"episode_id": 1, "event_time": 100,
+             "payload": {"value": 7}},
+        ])
+
+        rows = pmm.join_asof(
+            anchors.scan().select([
+                "episode_id", "event_time", "metadata.value"]),
+            secondary.scan().select("payload.value"),
+            on="event_time",
+            by="episode_id",
+            direction="nearest",
+            tolerance=0,
+        ).to_list()
+
+        self.assertEqual([1, None], [row["metadata_value"] for row in rows])
+        self.assertEqual(
+            [7, None], [row["payload_value"] for row in rows])
+
+        auth = TableQueryAuthResult(
+            filter=None,
+            column_masking={"payload": json.dumps({"name": "NULL"})},
+        )
+        selected = []
+
+        def query_auth(select):
+            selected.append(select)
+            return auth
+
+        secondary.raw_table.catalog_environment.table_query_auth = (
+            lambda options, identifier: query_auth)
+        masked = pmm.join_asof(
+            anchors.scan().select([
+                "episode_id", "event_time", "metadata.value"]),
+            secondary.scan().select("payload.value"),
+            on="event_time",
+            by="episode_id",
+            direction="nearest",
+            tolerance=0,
+        ).to_arrow()
+        self.assertEqual([None, None], masked["payload_value"].to_pylist())
+        self.assertTrue(masked.schema.field("payload_value").nullable)
+        self.assertIn(["payload"], selected)
+        self.assertTrue(all(
+            select is None or "payload_value" not in select
+            for select in selected
+        ))
+
+    def test_nested_projection_cannot_shadow_internal_row_id(self):
+        anchors = self._table("nested_row_id_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "_ROW": pa.struct([pa.field("ID", pa.int32())]),
+        })
+        samples = self._table("nested_row_id_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([
+            {"episode_id": 1, "event_time": 10, "_ROW": {"ID": 1}},
+            {"episode_id": 1, "event_time": 20, "_ROW": {"ID": 0}},
+        ])
+        samples.add([
+            {"episode_id": 1, "event_time": 10, "value": 7},
+            {"episode_id": 1, "event_time": 20, "value": 9},
+        ])
+
+        rows = pmm.join_asof(
+            anchors.scan().select(["_ROW.ID", "event_time"]),
+            samples.scan().select("value"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        ).to_list()
+
+        self.assertEqual([1, 0], [row["_ROW_ID"] for row in rows])
+        self.assertEqual([7, 9], [row["value"] for row in rows])
+
+    def test_nested_payload_name_cannot_shadow_group_key(self):
+        anchors = self._table("nested_key_anchors", {
+            "payload_value": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        samples = self._table("nested_key_samples", {
+            "payload_value": pa.int32(),
+            "event_time": pa.int64(),
+            "payload": pa.struct([pa.field("value", pa.int32())]),
+        })
+        anchors.add([{"payload_value": 1, "event_time": 10}])
+        samples.add([{
+            "payload_value": 1,
+            "event_time": 10,
+            "payload": {"value": 7},
+        }])
+
+        nested_only = pmm.join_asof(
+            anchors.scan(), samples.scan().select("payload.value"),
+            on="event_time", by="payload_value",
+            direction="nearest", tolerance=0,
+        ).to_list()[0]
+        self.assertEqual(7, nested_only["payload_value_right"])
+
+        both = pmm.join_asof(
+            anchors.scan(),
+            samples.scan().select(["payload.value", "payload_value"]),
+            on="event_time", by="payload_value",
+            direction="nearest", tolerance=0,
+        ).to_list()[0]
+        self.assertEqual(7, both["payload_value_right"])
+        self.assertNotIn("payload_value__0", both)
+
+    def test_alignment_rejects_unbound_nested_projection_masks(self):
+        anchors = self._table("nested_mask_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        source = self._table("nested_mask_source", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "payload": pa.struct([pa.field("value", pa.int32())]),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        source.add([{
+            "episode_id": 1,
+            "event_time": 100,
+            "payload": {"value": 7},
+        }])
+        auth = TableQueryAuthResult(
+            filter=None,
+            column_masking={"payload_value": json.dumps({"name": "NULL"})},
+        )
+        source.raw_table.catalog_environment.table_query_auth = (
+            lambda options, identifier: lambda select: auth)
+
+        with self.assertRaisesRegex(ValueError, "nested projection"):
+            pmm.join_asof(
+                anchors.scan(), source.scan().select("payload.value"),
+                on="event_time", by="episode_id",
+                direction="nearest", tolerance=0,
+            ).to_list()
+
+    def test_alignment_reuses_payload_scan_plans_across_batches(self):
+        anchors = self._table("planned_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        secondary = self._table("planned_secondary", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([
+            {"episode_id": 1, "event_time": value}
+            for value in range(8)
+        ])
+        secondary.add([
+            {"episode_id": 1, "event_time": value, "value": value}
+            for value in range(8)
+        ])
+        aligned = pmm.join_asof(
+            anchors.scan(), secondary.scan(),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        )
+        original_scan = FileScanner.scan
+
+        with mock.patch.object(
+                FileScanner, "scan", autospec=True,
+                side_effect=original_scan) as scan:
+            reader = aligned.to_arrow_batch_reader(batch_size=1)
+            self.assertEqual(8, sum(batch.num_rows for batch in reader))
+
+        self.assertEqual(4, scan.call_count)
+
+    def test_empty_source_stays_pinned_after_first_append(self):
+        anchors = self._table("pinned_empty_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        samples = self._table("pinned_empty_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 10}])
+        aligned = pmm.join_asof(
+            anchors.scan(), samples.scan().select("value"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        )
+
+        samples.add([{
+            "episode_id": 1, "event_time": 10, "value": 7,
+        }])
+
+        self.assertIsNone(aligned.resolved_snapshots["right_1"]["snapshot_id"])
+        self.assertIsNone(aligned.to_list()[0]["value"])
+
+    def test_alignment_reuses_decoded_parquet_row_groups_across_batches(self):
+        anchors = self._table("cached_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        samples = self._table("cached_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        row_count = 8192
+        rows = [
+            {"episode_id": 1, "event_time": value}
+            for value in range(row_count)
+        ]
+        anchors.add(rows)
+        samples.add([
+            dict(row, value=row["event_time"])
+            for row in rows
+        ])
+        original = FormatPyArrowReader._read_parquet_row_group_batches
+
+        for batch_size in (128, 1024, row_count):
+            with self.subTest(batch_size=batch_size):
+                decoded_rows = []
+
+                def tracked(reader, row_group, columns):
+                    for batch in original(reader, row_group, columns):
+                        if "value" in reader.existing_fields:
+                            decoded_rows.append(batch.num_rows)
+                        yield batch
+
+                with mock.patch.object(
+                        FormatPyArrowReader,
+                        "_read_parquet_row_group_batches", tracked):
+                    aligned = pmm.join_asof(
+                        anchors.scan(), samples.scan().select("value"),
+                        on="event_time", by="episode_id",
+                        direction="nearest", tolerance=0,
+                    )
+                    reader = aligned.to_arrow_batch_reader(
+                        batch_size=batch_size)
+                    result = pa.Table.from_batches(list(reader))
+
+                self.assertEqual(
+                    list(range(row_count)), result["value"].to_pylist())
+                self.assertEqual(row_count, sum(decoded_rows))
+
+    def test_alignment_streams_anchor_metadata(self):
+        anchors = self._table("streamed_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        samples = self._table("streamed_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([
+            {"episode_id": 1, "event_time": value}
+            for value in range(4)
+        ])
+        samples.add([
+            {"episode_id": 1, "event_time": value, "value": value}
+            for value in range(4)
+        ])
+        aligned = pmm.join_asof(
+            anchors.scan(), samples.scan(),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=0,
+        )
+
+        with mock.patch.object(
+                temporal, "_metadata_table",
+                wraps=temporal._metadata_table) as metadata_table:
+            self.assertEqual(4, len(aligned.to_list()))
+
+        # Only the indexed right side is collected into one metadata table.
+        self.assertEqual(1, metadata_table.call_count)
+
+    def test_alignment_splits_batches_before_arrow_offset_overflow(self):
+        child_count = 1 << 30
+        chunk = pa.ListArray.from_arrays(
+            pa.array([0, child_count], type=pa.int32()),
+            pa.nulls(child_count),
+        )
+        schema = pa.schema([pa.field("payload", chunk.type)])
+
+        class Fetcher:
+            def fetch(self, row_ids):
+                return pa.Table.from_arrays([
+                    pa.chunked_array([chunk for _ in row_ids])
+                ], schema=schema)
+
+        aligned = object.__new__(temporal.AsOfJoin)
+        aligned._anchor_schema = schema
+        aligned._sources = ()
+        rows = [{temporal._ROW_ID: value} for value in range(2)]
+
+        batches = list(aligned._build_batches(
+            rows, Fetcher(), [], schema))
+
+        self.assertEqual([1, 1], [batch.num_rows for batch in batches])
+        self.assertTrue(all(batch.validate() is None for batch in batches))
+        self.assertEqual(schema, batches[0].schema)
+
+    def test_alignment_closes_anchor_stream_when_reader_closes(self):
+        anchors = self._table("closable_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        samples = self._table("closable_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "value": pa.int32(),
+        })
+        anchors.add([
+            {"episode_id": 1, "event_time": value}
+            for value in range(2)
+        ])
+        samples.add([
+            {"episode_id": 1, "event_time": value, "value": value}
+            for value in range(2)
+        ])
+        closed = []
+        original = temporal._metadata_batches
+
+        def tracked_batches(*args):
+            try:
+                for batch in original(*args):
+                    yield batch
+            finally:
+                closed.append(True)
+
+        with mock.patch.object(
+                temporal, "_metadata_batches", tracked_batches):
+            reader = pmm.join_asof(
+                anchors.scan(), samples.scan().select("value"),
+                on="event_time", by="episode_id",
+                direction="nearest", tolerance=0,
+            ).to_arrow_batch_reader(batch_size=1)
+            next(reader)
+            reader.close()
+
+        self.assertEqual([True], closed)
+
+    def test_alignment_keeps_blob_payloads_as_descriptors(self):
+        anchors = self._table("blob_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+        })
+        images = self._table("blob_images", {
+            "episode_id": pa.int32(),
+            "event_time": pa.int64(),
+            "image": pa.large_binary(),
+        })
+        anchors.add([{"episode_id": 1, "event_time": 100}])
+        images.add([{
+            "episode_id": 1,
+            "event_time": 100,
+            "image": b"encoded-image",
+        }])
+
+        row = pmm.join_asof(
+            anchors.scan(),
+            images.scan().select("image"),
+            on="event_time",
+            by="episode_id",
+            direction="nearest",
+            tolerance=0,
+        ).to_list()[0]
+
+        descriptor = pmm.BlobDescriptor.deserialize(row["image"])
+        self.assertTrue(descriptor.uri.endswith(".blob"))
+        self.assertEqual(len(b"encoded-image"), descriptor.length)
+
+    def test_alignment_supports_timestamp_columns_with_different_names(self):
+        anchors = self._table("timestamp_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.timestamp("ms"),
+        })
+        samples = self._table("timestamp_samples", {
+            "episode_id": pa.int32(),
+            "captured_at": pa.timestamp("ms"),
+            "value": pa.int32(),
+        })
+        anchor_time = datetime(2026, 9, 1, 12, 0, 0, 100000)
+        sample_time = anchor_time - timedelta(milliseconds=5)
+        anchors.add([{"episode_id": 1, "event_time": anchor_time}])
+        samples.add([{
+            "episode_id": 1,
+            "captured_at": sample_time,
+            "value": 7,
+        }])
+
+        row = pmm.join_asof(
+            anchors.scan(),
+            samples.scan().select("value"),
+            on="event_time",
+            by="episode_id",
+            direction="nearest",
+            right_on="captured_at",
+            tolerance=timedelta(milliseconds=10),
+        ).to_list()[0]
+
+        self.assertEqual(7, row["value"])
+
+    def test_alignment_preserves_nanosecond_timestamp_precision(self):
+        anchors = self._table("nanosecond_anchors", {
+            "episode_id": pa.int32(),
+            "event_time": pa.timestamp("ns"),
+        })
+        samples = self._table("nanosecond_samples", {
+            "episode_id": pa.int32(),
+            "event_time": pa.timestamp("ns"),
+            "value": pa.int32(),
+        })
+        anchors.add(pa.table({
+            "episode_id": pa.array([1], type=pa.int32()),
+            "event_time": pa.array(
+                [1_000_000_001], type=pa.int64()).cast(pa.timestamp("ns")),
+        }))
+        samples.add(pa.table({
+            "episode_id": pa.array([1, 1], type=pa.int32()),
+            "event_time": pa.array(
+                [1_000_000_000, 1_000_000_001],
+                type=pa.int64()).cast(pa.timestamp("ns")),
+            "value": pa.array([1, 2], type=pa.int32()),
+        }))
+
+        row = pmm.join_asof(
+            anchors.scan(), samples.scan().select("value"),
+            on="event_time", by="episode_id",
+            direction="nearest", tolerance=timedelta(0),
+        ).to_list()[0]
+
+        self.assertEqual(2, row["value"])
+
+    def _table(self, name, fields):
+        return self.conn.create_table(name, schema=pa.schema([
+            pa.field(field_name, field_type)
+            for field_name, field_type in fields.items()
+        ]))
+
+
+if __name__ == "__main__":
+    unittest.main()
diff --git a/paimon-python/pypaimon/tests/parquet_metadata_cache_test.py 
b/paimon-python/pypaimon/tests/parquet_metadata_cache_test.py
index 2ca91f5776..257560ef63 100644
--- a/paimon-python/pypaimon/tests/parquet_metadata_cache_test.py
+++ b/paimon-python/pypaimon/tests/parquet_metadata_cache_test.py
@@ -230,7 +230,9 @@ class FileFormatMetadataCacheTest(unittest.TestCase):
     def test_does_not_retain_entry_larger_than_size_limit(self):
         cache = reader_module._FileFormatDatasetCache(5)
         loads = []
+        small_key = (None, "parquet", "small")
         key = (None, "parquet", "large")
+        cache.get_or_load(small_key, lambda: "small", lambda _: 4)
 
         def load():
             loads.append(True)
@@ -241,8 +243,8 @@ class FileFormatMetadataCacheTest(unittest.TestCase):
         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)
+        self.assertEqual([small_key], list(cache._entries))
+        self.assertEqual(4, cache.estimated_size)
 
     def test_does_not_retain_entry_without_size_estimate(self):
         cache = reader_module._FileFormatDatasetCache(10)
diff --git a/paimon-python/pypaimon/tests/parquet_row_range_test.py 
b/paimon-python/pypaimon/tests/parquet_row_range_test.py
index 6bba358924..5318e3d0e1 100644
--- a/paimon-python/pypaimon/tests/parquet_row_range_test.py
+++ b/paimon-python/pypaimon/tests/parquet_row_range_test.py
@@ -27,7 +27,10 @@ import pyarrow.fs as pafs
 import pyarrow.parquet as pq
 
 from pypaimon import CatalogFactory, Schema
-from pypaimon.read.reader.format_pyarrow_reader import FormatPyArrowReader
+from pypaimon.read.reader.format_pyarrow_reader import (
+    FormatPyArrowReader,
+    _DecodedRowGroupCache,
+)
 from pypaimon.schema.data_types import AtomicType, DataField
 
 
@@ -388,6 +391,55 @@ class ParquetRowRangeTest(unittest.TestCase):
             ],
         )
 
+    def test_oversized_row_group_bypasses_decoded_cache(self):
+        path = os.path.join(self.tempdir, "oversized-row-group.parquet")
+        pq.write_table(
+            pa.table({"payload": [b"x" * 256] * 16}),
+            path,
+            row_group_size=16,
+            compression="none",
+        )
+        cache = _DecodedRowGroupCache(1024)
+        small_key = ("small",)
+        list(cache.iter_or_load(
+            small_key,
+            lambda: iter([pa.record_batch({"value": pa.array([1])})]),
+        ))
+        decoded_rows = []
+        original = FormatPyArrowReader._read_parquet_row_group_batches
+
+        def tracked(reader, row_group, columns):
+            for batch in original(reader, row_group, columns):
+                decoded_rows.append(batch.num_rows)
+                yield batch
+
+        with mock.patch.object(
+                FormatPyArrowReader,
+                "_read_parquet_row_group_batches", tracked):
+            reader = FormatPyArrowReader(
+                _LocalFileIO(),
+                "parquet",
+                path,
+                [DataField(0, "payload", AtomicType("BYTES"))],
+                None,
+                batch_size=2,
+                row_ranges=[(0, 0)],
+                row_group_cache=cache,
+            )
+            try:
+                self.assertEqual(
+                    [b"x" * 256],
+                    reader.read_arrow_batch().column(0).to_pylist(),
+                )
+                self.assertEqual(2, sum(decoded_rows))
+                while reader.read_arrow_batch() is not None:
+                    pass
+            finally:
+                reader.close()
+
+        self.assertEqual(16, sum(decoded_rows))
+        self.assertEqual([small_key], list(cache._cache._entries))
+
 
 if __name__ == "__main__":
     unittest.main()

Reply via email to