From 23656de045d636a0b8d0b9cca57ec620a2fcc7a8 Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Mon, 18 May 2026 18:07:25 -0700 Subject: [PATCH 01/14] [Data] Add task-based upsert for iceberg using scan merge approach Signed-off-by: Ayush Kumar --- .../_internal/datasource/iceberg_datasink.py | 274 ++++++++++++++++-- 1 file changed, 245 insertions(+), 29 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index fbfd92919905..7e5b75f63d05 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -2,6 +2,7 @@ Module to write a Ray Dataset into an iceberg table, by using the Ray Datasink API. """ import logging +import ray from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional, Union @@ -20,13 +21,109 @@ from pyiceberg.io import FileIO from pyiceberg.manifest import DataFile from pyiceberg.schema import Schema - from pyiceberg.table import Table + from pyiceberg.table import FileScanTask, Table from pyiceberg.table.metadata import TableMetadata from pyiceberg.table.update.schema import UpdateSchema logger = logging.getLogger(__name__) +@ray.remote +def _rewrite_iceberg_file( + file_scan_task: "FileScanTask", + keys_ref: "pa.Table", + upsert_cols: List[str], + table_metadata: "TableMetadata", + io: "FileIO", +) -> "tuple[Optional[DataFile], List[DataFile]]": + """Read one Iceberg file, anti-join against upsert keys, write false-positive rows. + + False positives are rows in the file that are NOT in the upsert batch — the + coarse range filter would delete them, so we preserve them by writing them + as new data files before the delete. + + Returns (original DataFile to delete, list of new FP DataFiles). + If the entire file is matched (no FPs), returns (file, []). + If the file has no matched rows at all, returns (None, []) — leave it untouched. + """ + import hashlib + import time as _time + import uuid as _uuid + + import pyarrow as pa + from pyiceberg.expressions import AlwaysTrue + from pyiceberg.io.pyarrow import ArrowScan, _dataframe_to_data_files + + file_path = file_scan_task.file.file_path + file_size_mb = file_scan_task.file.file_size_in_bytes / 1e6 + t_start = _time.perf_counter() + + batch = ArrowScan( + table_metadata=table_metadata, + io=io, + projected_schema=table_metadata.schema(), + row_filter=AlwaysTrue(), + ).to_table(tasks=[file_scan_task]) + + t_read = _time.perf_counter() + logger.debug( + "[rewrite] read %d rows / %.1f MB (compressed) from %s in %.2fs", + len(batch), + file_size_mb, + file_path.split("/")[-1], + t_read - t_start, + ) + + if len(batch) == 0: + return (None, []) + + # Cast batch key columns to match keys_ref types so PyArrow's join doesn't + # raise ArrowInvalid on utf8/large_utf8 or similar width mismatches. + key_cast = {f.name: f.type for f in keys_ref.schema if f.name in upsert_cols} + batch_keys = batch.select(upsert_cols) + for col_name, target_type in key_cast.items(): + if batch_keys.schema.field(col_name).type != target_type: + col_idx = batch_keys.schema.get_field_index(col_name) + batch_keys = batch_keys.set_column( + col_idx, col_name, batch_keys[col_name].cast(target_type) + ) + + idx_col = pa.array(range(len(batch)), type=pa.int64()) + fp_keys = ( + batch_keys.append_column("__row_idx__", idx_col) + .join(keys_ref, keys=upsert_cols, join_type="left anti") + ) + + if len(fp_keys) == 0: + # Every row in this file is being upserted — delete the whole file, no FP file needed. + logger.debug("[rewrite] %s: all %d rows matched -> whole-file delete", file_path.split("/")[-1], len(batch)) + return (file_scan_task.file, []) + + if len(fp_keys) == len(batch): + # No rows in this file match any upsert key — leave it alone entirely. + logger.debug("[rewrite] %s: 0 rows matched -> untouched", file_path.split("/")[-1]) + return (None, []) + + fp_rows = batch.take(fp_keys["__row_idx__"]) + # Derive a deterministic write_uuid from the source file path so that + # task retries overwrite the same object rather than leaking orphan files. + fp_write_uuid = _uuid.UUID(hashlib.md5(file_path.encode()).hexdigest()) + fp_files = list( + _dataframe_to_data_files( + table_metadata=table_metadata, df=fp_rows, io=io, write_uuid=fp_write_uuid + ) + ) + logger.debug( + "[rewrite] %s: %d/%d rows are FPs -> wrote %d FP file(s) in %.2fs", + file_path.split("/")[-1], + len(fp_keys), + len(batch), + len(fp_files), + _time.perf_counter() - t_read, + ) + return (file_scan_task.file, fp_files) + + @dataclass class IcebergWriteResult: """Result from writing blocks to Iceberg storage. @@ -205,6 +302,149 @@ def _get_upsert_cols(self) -> List[str]: upsert_cols.append(col_name) return upsert_cols + def _build_coarse_range_filter( + self, + keys_table: "pa.Table", + upsert_cols: List[str], + ) -> "BooleanExpression": + """Build an O(1) coarse range filter covering all upsert key values. + + For each upsert column computes AND(GTE(col, min), LTE(col, max)). + The filter may match rows outside the upsert batch (false positives); + callers must anti-join to identify and preserve those rows. + """ + import pyarrow.compute as pc + from pyiceberg.expressions import AlwaysTrue, And, GreaterThanOrEqual, LessThanOrEqual + + expr = None + for col_name in upsert_cols: + mm = pc.min_max(keys_table[col_name]) + min_val = mm["min"].as_py() + max_val = mm["max"].as_py() + if min_val is None: + continue + col_expr = And( + GreaterThanOrEqual(col_name, min_val), + LessThanOrEqual(col_name, max_val), + ) + expr = col_expr if expr is None else And(expr, col_expr) + + return expr if expr is not None else AlwaysTrue() + + def _commit_upsert_scan_merge( + self, + txn: "Table.transaction", + data_files: List["DataFile"], + keys_table: "pa.Table", + upsert_cols: List[str], + delete_kwargs: Dict[str, Any], + ) -> None: + """Upsert commit using coarse range filter + per-file distributed anti-join. + + 1. Build an O(1) coarse range filter covering all upsert key values. + 2. plan_files() on the driver — manifest reads only, no data I/O on driver. + 3. Dispatch one Ray task per candidate file. Each task reads its file, + anti-joins against the upsert keys to find false positives (rows that + the coarse delete would remove but that are NOT being upserted), and + writes them as new data files directly to storage. + 4. Commit atomically via txn.update_snapshot().overwrite(): delete each + original candidate file and append FP files + new upsert data files. + + No table data ever flows through the driver process, avoiding driver + OOM on wide-schema tables. + """ + import time + + import pyarrow as pa + + # Dedup keys to minimise per-task anti-join hash table size. + keys_table = keys_table.group_by(upsert_cols).aggregate([]) + + coarse_filter = self._build_coarse_range_filter(keys_table, upsert_cols) + logger.debug("[scan-merge] coarse_filter=%s", coarse_filter) + + # plan_files() reads only manifest metadata — no Parquet data on the driver. + t0 = time.perf_counter() + file_scan_tasks = list(self._table.scan(row_filter=coarse_filter).plan_files()) + logger.info( + "[scan-merge] planned %d candidate file(s) in %.2fs", + len(file_scan_tasks), + time.perf_counter() - t0, + ) + + if not file_scan_tasks: + # No existing files match the coarse filter — pure inserts only. + self._append_and_commit(txn, data_files) + return + + # Put the deduped keys in the object store once; all tasks share one copy. + keys_ref = ray.put(keys_table) + + # Parquet decompression factor: compressed Parquet typically expands + # ~8-10x in memory once decoded. Conservative default keeps Ray's + # scheduler from stacking too many rewrite tasks on one node. + _PARQUET_EXPANSION = 8 + + t0 = time.perf_counter() + refs = [ + _rewrite_iceberg_file.options( + memory=int(task.file.file_size_in_bytes * _PARQUET_EXPANSION) + ).remote(task, keys_ref, upsert_cols, self._table_metadata, self._io) + for task in file_scan_tasks + ] + logger.info("[scan-merge] dispatched %d rewrite task(s)", len(refs)) + + # Collect results with periodic progress logs so long rewrites aren't silent. + results = [] + pending = list(refs) + _LOG_INTERVAL = max(1, len(refs) // 10) # log ~10 times total + while pending: + done, pending = ray.wait(pending, num_returns=min(_LOG_INTERVAL, len(pending))) + results.extend(ray.get(done)) + logger.debug( + "[scan-merge] rewrite progress: %d/%d file(s) done (%.1fs elapsed)", + len(results), + len(refs), + time.perf_counter() - t0, + ) + + logger.info( + "[scan-merge] all %d file(s) rewritten in %.2fs", + len(refs), + time.perf_counter() - t0, + ) + + # Count how many files were wholly deleted vs partially rewritten. + n_whole_delete = sum(1 for old, fps in results if old is not None and not fps) + n_partial = sum(1 for old, fps in results if fps) + n_untouched = sum(1 for old, fps in results if old is None) + logger.info( + "[scan-merge] files: %d whole-delete, %d partial-rewrite, %d untouched", + n_whole_delete, + n_partial, + n_untouched, + ) + + # Single atomic commit: schema update (already staged in txn) + this overwrite. + # _OverwriteFiles handles both file-level deletes and appends in one snapshot. + t0 = time.perf_counter() + with txn.update_snapshot( + snapshot_properties=self._snapshot_properties + ).overwrite() as snap: + for old_file, fp_files in results: + if old_file is not None: + snap.delete_data_file(old_file) + for fp_file in fp_files: + snap.append_data_file(fp_file) + for df in data_files: + snap.append_data_file(df) + + self._with_retry( + txn.commit_transaction, + description=f"commit upsert transaction to Iceberg table '{self.table_identifier}'", + ) + logger.info("[scan-merge] committed in %.2fs", time.perf_counter() - t0) + def _append_and_commit( self, txn: "Table.transaction", data_files: List["DataFile"] ) -> None: @@ -241,7 +481,6 @@ def _commit_upsert( import time import pyarrow as pa - from pyiceberg.table.upsert_util import create_match_filter # Create delete filter if we have join keys if upsert_keys is not None and len(upsert_keys) > 0: @@ -267,39 +506,16 @@ def _commit_upsert( # Only delete if we have non-NULL keys if len(keys_table) > 0: - logger.info( - "[upsert commit] Building delete filter from %d keys (cols: %s) ...", - len(keys_table), - upsert_cols, - ) - t0 = time.perf_counter() - # Use PyIceberg's helper to build delete filter - delete_filter = create_match_filter(keys_table, upsert_cols) - logger.info( - "[upsert commit] create_match_filter done in %.2fs: filter type=%s", - time.perf_counter() - t0, - type(delete_filter).__name__, - ) - - # Prepare kwargs for delete delete_kwargs = self._upsert_kwargs.copy() delete_kwargs.pop(_UPSERT_COLS_ID, None) - - logger.info("[upsert commit] Executing txn.delete() ...") - t0 = time.perf_counter() - txn.delete( - delete_filter=delete_filter, - snapshot_properties=self._snapshot_properties, - **delete_kwargs, - ) - logger.info( - "[upsert commit] txn.delete() done in %.2fs", - time.perf_counter() - t0, + self._commit_upsert_scan_merge( + txn, data_files, keys_table, upsert_cols, delete_kwargs ) + return else: logger.info("[upsert commit] No upsert keys — skipping delete phase") - # Append new data files (includes updates and inserts) and commit + # No non-NULL keys — just append new data files and commit logger.info( "[upsert commit] Appending %d data files and committing ...", len(data_files), From d88ad8e53832efa3b673c44f475aa7261911705f Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Mon, 18 May 2026 18:38:30 -0700 Subject: [PATCH 02/14] [Data] Add tests for iceberg upsert using scan merge Signed-off-by: Ayush Kumar --- .../ray/data/tests/datasource/test_iceberg.py | 260 ++++++++++++++++++ 1 file changed, 260 insertions(+) diff --git a/python/ray/data/tests/datasource/test_iceberg.py b/python/ray/data/tests/datasource/test_iceberg.py index 2bff444b7576..f04f535a7bec 100644 --- a/python/ray/data/tests/datasource/test_iceberg.py +++ b/python/ray/data/tests/datasource/test_iceberg.py @@ -1301,6 +1301,266 @@ def test_write_upsert_empty_table(self, clean_table): assert rows_same(result, expected) +@pytest.mark.skipif( + get_pyarrow_version() < parse_version("14.0.0"), + reason="PyIceberg 0.7.0 fails on pyarrow <= 14.0.0", +) +class TestUpsertScanMerge: + """Test the scan-merge upsert algorithm for correctness. + + See ``IcebergDatasink._commit_upsert_scan_merge`` for algorithm details. + """ + + def test_upsert_preserves_false_positives_sparse_keys(self, clean_table): + """Sparse upsert keys leave intermediate rows as false positives that + must be preserved after the rewrite.""" + from ray.data import SaveMode + + seed = _create_typed_dataframe( + { + "col_a": list(range(1, 11)), + "col_b": [f"seed_{i}" for i in range(1, 11)], + "col_c": [1] * 10, + } + ) + _write_to_iceberg(seed) + + upsert_data = _create_typed_dataframe( + { + "col_a": [1, 10], + "col_b": ["updated_1", "updated_10"], + "col_c": [1, 1], + } + ) + _write_to_iceberg( + upsert_data, + mode=SaveMode.UPSERT, + upsert_kwargs={"join_cols": ["col_a"]}, + ) + + result = _read_from_iceberg(sort_by="col_a") + expected = _create_typed_dataframe( + { + "col_a": list(range(1, 11)), + "col_b": ["updated_1"] + + [f"seed_{i}" for i in range(2, 10)] + + ["updated_10"], + "col_c": [1] * 10, + } + ) + assert rows_same(result, expected) + + def test_upsert_across_multiple_files(self, clean_table): + """Two separate seed writes produce at least two data files. A sparse + upsert that spans both files must rewrite false positives in each.""" + from ray.data import SaveMode + + _write_to_iceberg( + _create_typed_dataframe( + { + "col_a": [1, 2, 3], + "col_b": ["seed_1", "seed_2", "seed_3"], + "col_c": [1, 1, 1], + } + ) + ) + _write_to_iceberg( + _create_typed_dataframe( + { + "col_a": [10, 11, 12], + "col_b": ["seed_10", "seed_11", "seed_12"], + "col_c": [1, 1, 1], + } + ) + ) + + upsert_data = _create_typed_dataframe( + { + "col_a": [1, 12], + "col_b": ["updated_1", "updated_12"], + "col_c": [1, 1], + } + ) + _write_to_iceberg( + upsert_data, + mode=SaveMode.UPSERT, + upsert_kwargs={"join_cols": ["col_a"]}, + ) + + result = _read_from_iceberg(sort_by="col_a") + expected = _create_typed_dataframe( + { + "col_a": [1, 2, 3, 10, 11, 12], + "col_b": [ + "updated_1", + "seed_2", + "seed_3", + "seed_10", + "seed_11", + "updated_12", + ], + "col_c": [1] * 6, + } + ) + assert rows_same(result, expected) + + def test_upsert_whole_file_delete_when_all_keys_match(self, clean_table): + """When every seed row in the coarse range is in the upsert batch, the + original file is wholly deleted and only upsert rows remain — no + duplicates.""" + from ray.data import SaveMode + + seed = _create_typed_dataframe( + { + "col_a": [1, 2, 3, 4, 5], + "col_b": ["seed_1", "seed_2", "seed_3", "seed_4", "seed_5"], + "col_c": [1] * 5, + } + ) + _write_to_iceberg(seed) + + upsert_data = _create_typed_dataframe( + { + "col_a": [1, 2, 3, 4, 5], + "col_b": [f"updated_{i}" for i in range(1, 6)], + "col_c": [1] * 5, + } + ) + _write_to_iceberg( + upsert_data, + mode=SaveMode.UPSERT, + upsert_kwargs={"join_cols": ["col_a"]}, + ) + + result = _read_from_iceberg(sort_by="col_a") + expected = _create_typed_dataframe( + { + "col_a": [1, 2, 3, 4, 5], + "col_b": [f"updated_{i}" for i in range(1, 6)], + "col_c": [1] * 5, + } + ) + assert rows_same(result, expected) + + def test_upsert_pure_insert_short_circuit(self, clean_table): + """Upsert keys outside the seed's coarse range hit zero candidate + files; the new rows must still be appended.""" + from ray.data import SaveMode + + seed = _create_typed_dataframe( + { + "col_a": [1, 2, 3], + "col_b": ["seed_1", "seed_2", "seed_3"], + "col_c": [1, 1, 1], + } + ) + _write_to_iceberg(seed) + + upsert_data = _create_typed_dataframe( + { + "col_a": [100, 101], + "col_b": ["new_100", "new_101"], + "col_c": [1, 1], + } + ) + _write_to_iceberg( + upsert_data, + mode=SaveMode.UPSERT, + upsert_kwargs={"join_cols": ["col_a"]}, + ) + + result = _read_from_iceberg(sort_by="col_a") + expected = _create_typed_dataframe( + { + "col_a": [1, 2, 3, 100, 101], + "col_b": ["seed_1", "seed_2", "seed_3", "new_100", "new_101"], + "col_c": [1] * 5, + } + ) + assert rows_same(result, expected) + + def test_upsert_composite_key_preserves_false_positives(self, clean_table): + """Composite-key anti-join must match on all join columns; rows that + share one column with an upsert key but not the full composite are + false positives that must be preserved.""" + from ray.data import SaveMode + + composites = [(a, b) for a in [1, 2, 3] for b in ["x", "y", "z"]] + seed = _create_typed_dataframe( + { + "col_a": [a for a, _ in composites], + "col_b": [b for _, b in composites], + "col_c": [1] * len(composites), + } + ) + _write_to_iceberg(seed) + + upsert_data = _create_typed_dataframe( + { + "col_a": [1, 3], + "col_b": ["x", "z"], + "col_c": [99, 99], + } + ) + _write_to_iceberg( + upsert_data, + mode=SaveMode.UPSERT, + upsert_kwargs={"join_cols": ["col_a", "col_b"]}, + ) + + result = _read_from_iceberg(sort_by=["col_a", "col_b"]) + expected_col_c = [ + 99 if (a, b) in {(1, "x"), (3, "z")} else 1 + for a, b in sorted(composites) + ] + expected = _create_typed_dataframe( + { + "col_a": [a for a, _ in sorted(composites)], + "col_b": [b for _, b in sorted(composites)], + "col_c": expected_col_c, + } + ) + assert rows_same(result, expected) + + def test_upsert_string_key(self, clean_table): + """String join column exercises the type-cast path in + _rewrite_iceberg_file that aligns utf8 / large_utf8 between the file + batch and the upsert-keys table.""" + from ray.data import SaveMode + + seed = _create_typed_dataframe( + { + "col_a": [1, 2, 3, 4, 5], + "col_b": ["a", "b", "c", "d", "e"], + "col_c": [1] * 5, + } + ) + _write_to_iceberg(seed) + + upsert_data = _create_typed_dataframe( + { + "col_a": [10, 20], + "col_b": ["a", "e"], + "col_c": [1, 1], + } + ) + _write_to_iceberg( + upsert_data, + mode=SaveMode.UPSERT, + upsert_kwargs={"join_cols": ["col_b"]}, + ) + + result = _read_from_iceberg(sort_by="col_b") + expected = _create_typed_dataframe( + { + "col_a": [10, 2, 3, 4, 20], + "col_b": ["a", "b", "c", "d", "e"], + "col_c": [1] * 5, + } + ) + assert rows_same(result, expected) + + @pytest.fixture def table_with_identifier_fields() -> Generator[Tuple[Catalog, Table], None, None]: """Pytest fixture to create a table with identifier fields for upsert tests.""" From 47fbe18adb80be34fc025746fca19332bb1615fd Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Mon, 18 May 2026 18:44:57 -0700 Subject: [PATCH 03/14] [Data] Add BooleanExpr import and lint fixes Signed-off-by: Ayush Kumar --- .../_internal/datasource/iceberg_datasink.py | 31 +++++++++++++------ .../ray/data/tests/datasource/test_iceberg.py | 7 ++--- 2 files changed, 24 insertions(+), 14 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index 7e5b75f63d05..6867af2c8750 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -2,10 +2,10 @@ Module to write a Ray Dataset into an iceberg table, by using the Ray Datasink API. """ import logging -import ray from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional, Union +import ray from ray._common.retry import call_with_retry from ray.data._internal.execution.interfaces import TaskContext from ray.data._internal.savemode import SaveMode @@ -18,6 +18,7 @@ if TYPE_CHECKING: import pyarrow as pa from pyiceberg.catalog import Catalog + from pyiceberg.expressions import BooleanExpression from pyiceberg.io import FileIO from pyiceberg.manifest import DataFile from pyiceberg.schema import Schema @@ -89,19 +90,24 @@ def _rewrite_iceberg_file( ) idx_col = pa.array(range(len(batch)), type=pa.int64()) - fp_keys = ( - batch_keys.append_column("__row_idx__", idx_col) - .join(keys_ref, keys=upsert_cols, join_type="left anti") + fp_keys = batch_keys.append_column("__row_idx__", idx_col).join( + keys_ref, keys=upsert_cols, join_type="left anti" ) if len(fp_keys) == 0: # Every row in this file is being upserted — delete the whole file, no FP file needed. - logger.debug("[rewrite] %s: all %d rows matched -> whole-file delete", file_path.split("/")[-1], len(batch)) + logger.debug( + "[rewrite] %s: all %d rows matched -> whole-file delete", + file_path.split("/")[-1], + len(batch), + ) return (file_scan_task.file, []) if len(fp_keys) == len(batch): # No rows in this file match any upsert key — leave it alone entirely. - logger.debug("[rewrite] %s: 0 rows matched -> untouched", file_path.split("/")[-1]) + logger.debug( + "[rewrite] %s: 0 rows matched -> untouched", file_path.split("/")[-1] + ) return (None, []) fp_rows = batch.take(fp_keys["__row_idx__"]) @@ -314,7 +320,12 @@ def _build_coarse_range_filter( callers must anti-join to identify and preserve those rows. """ import pyarrow.compute as pc - from pyiceberg.expressions import AlwaysTrue, And, GreaterThanOrEqual, LessThanOrEqual + from pyiceberg.expressions import ( + AlwaysTrue, + And, + GreaterThanOrEqual, + LessThanOrEqual, + ) expr = None for col_name in upsert_cols: @@ -355,8 +366,6 @@ def _commit_upsert_scan_merge( """ import time - import pyarrow as pa - # Dedup keys to minimise per-task anti-join hash table size. keys_table = keys_table.group_by(upsert_cols).aggregate([]) @@ -399,7 +408,9 @@ def _commit_upsert_scan_merge( pending = list(refs) _LOG_INTERVAL = max(1, len(refs) // 10) # log ~10 times total while pending: - done, pending = ray.wait(pending, num_returns=min(_LOG_INTERVAL, len(pending))) + done, pending = ray.wait( + pending, num_returns=min(_LOG_INTERVAL, len(pending)) + ) results.extend(ray.get(done)) logger.debug( "[scan-merge] rewrite progress: %d/%d file(s) done (%.1fs elapsed)", diff --git a/python/ray/data/tests/datasource/test_iceberg.py b/python/ray/data/tests/datasource/test_iceberg.py index f04f535a7bec..5f126b82848f 100644 --- a/python/ray/data/tests/datasource/test_iceberg.py +++ b/python/ray/data/tests/datasource/test_iceberg.py @@ -1308,8 +1308,8 @@ def test_write_upsert_empty_table(self, clean_table): class TestUpsertScanMerge: """Test the scan-merge upsert algorithm for correctness. - See ``IcebergDatasink._commit_upsert_scan_merge`` for algorithm details. - """ + See ``IcebergDatasink._commit_upsert_scan_merge`` for algorithm details. + """ def test_upsert_preserves_false_positives_sparse_keys(self, clean_table): """Sparse upsert keys leave intermediate rows as false positives that @@ -1510,8 +1510,7 @@ def test_upsert_composite_key_preserves_false_positives(self, clean_table): result = _read_from_iceberg(sort_by=["col_a", "col_b"]) expected_col_c = [ - 99 if (a, b) in {(1, "x"), (3, "z")} else 1 - for a, b in sorted(composites) + 99 if (a, b) in {(1, "x"), (3, "z")} else 1 for a, b in sorted(composites) ] expected = _create_typed_dataframe( { From 1a21b793b8441d09bc00784766fd5ecd9623970e Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Mon, 18 May 2026 18:51:59 -0700 Subject: [PATCH 04/14] [Data] Modify docstring for iceberg upsert rewrite task and algorithm Signed-off-by: Ayush Kumar --- .../ray/data/_internal/datasource/iceberg_datasink.py | 11 ++++------- 1 file changed, 4 insertions(+), 7 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index 6867af2c8750..f8db91cc8ec9 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -39,8 +39,8 @@ def _rewrite_iceberg_file( ) -> "tuple[Optional[DataFile], List[DataFile]]": """Read one Iceberg file, anti-join against upsert keys, write false-positive rows. - False positives are rows in the file that are NOT in the upsert batch — the - coarse range filter would delete them, so we preserve them by writing them + False positives are rows in the file that are not in the upsert batch. The + coarse range filter (see ``IcebergDatasink._build_coarse_range_filter``) would delete them, so we preserve them by writing them as new data files before the delete. Returns (original DataFile to delete, list of new FP DataFiles). @@ -352,17 +352,14 @@ def _commit_upsert_scan_merge( ) -> None: """Upsert commit using coarse range filter + per-file distributed anti-join. - 1. Build an O(1) coarse range filter covering all upsert key values. - 2. plan_files() on the driver — manifest reads only, no data I/O on driver. + 1. Build an O(1) coarse range filter using min-max covering upsert key values (for each column). + 2. plan_files() on the driver to find candidate files that could be updated 3. Dispatch one Ray task per candidate file. Each task reads its file, anti-joins against the upsert keys to find false positives (rows that the coarse delete would remove but that are NOT being upserted), and writes them as new data files directly to storage. 4. Commit atomically via txn.update_snapshot().overwrite(): delete each original candidate file and append FP files + new upsert data files. - - No table data ever flows through the driver process, avoiding driver - OOM on wide-schema tables. """ import time From 73b254e10112b1c0ebf39618042f3f4afd526d33 Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Tue, 19 May 2026 14:51:06 -0700 Subject: [PATCH 05/14] [Data] Refactor upsert code to wire through upsert_kwargs, _append_and_commit() Signed-off-by: Ayush Kumar --- .../_internal/datasource/iceberg_datasink.py | 77 ++++++++++++++----- 1 file changed, 59 insertions(+), 18 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index f8db91cc8ec9..999f2a363464 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -40,12 +40,12 @@ def _rewrite_iceberg_file( """Read one Iceberg file, anti-join against upsert keys, write false-positive rows. False positives are rows in the file that are not in the upsert batch. The - coarse range filter (see ``IcebergDatasink._build_coarse_range_filter``) would delete them, so we preserve them by writing them - as new data files before the delete. + coarse range filter would delete them (see ``IcebergDatasink._build_coarse_range_filter``), + so we preserve them by writing them as new data files before the delete. Returns (original DataFile to delete, list of new FP DataFiles). If the entire file is matched (no FPs), returns (file, []). - If the file has no matched rows at all, returns (None, []) — leave it untouched. + If the file has no matched rows at all, returns (None, []), leave it untouched. """ import hashlib import time as _time @@ -301,11 +301,34 @@ def _get_upsert_cols(self) -> List[str]: upsert_cols = self._upsert_kwargs.get(_UPSERT_COLS_ID, []) if not upsert_cols: # Use table's identifier fields as fallback + identifier_cols = [] schema = self._table_metadata.schema() for field_id in schema.identifier_field_ids: col_name = schema.find_column_name(field_id) if col_name: - upsert_cols.append(col_name) + identifier_cols.append(col_name) + return identifier_cols + + case_sensitive = self._upsert_kwargs.get("case_sensitive", True) + + # To support case insensitivity, we need to define a mapping of + # provided (possibly case-modified) names to their original names in the schema + if not case_sensitive: + schema = self._table_metadata.schema() + lower_to_original_mapping = { + col.name.lower(): col.name for col in schema.fields + } + resolved_upsert_cols = [] + for upsert_col in upsert_cols: + resolved_col = lower_to_original_mapping.get(upsert_col.lower()) + if resolved_col is None: + raise ValueError( + f"Upsert join column {upsert_col!r} does not match any column in " + f"table schema (case-insensitive)." + ) + resolved_upsert_cols.append(resolved_col) + upsert_cols = resolved_upsert_cols + return upsert_cols def _build_coarse_range_filter( @@ -348,7 +371,6 @@ def _commit_upsert_scan_merge( data_files: List["DataFile"], keys_table: "pa.Table", upsert_cols: List[str], - delete_kwargs: Dict[str, Any], ) -> None: """Upsert commit using coarse range filter + per-file distributed anti-join. @@ -363,6 +385,18 @@ def _commit_upsert_scan_merge( """ import time + case_sensitive = self._upsert_kwargs.get("case_sensitive", True) + branch = self._upsert_kwargs.get("branch", "main") + unknown = set(self._upsert_kwargs) - { + _UPSERT_COLS_ID, + "case_sensitive", + "branch", + } + if unknown: + logger.warning( + "[scan-merge] ignoring unsupported upsert_kwargs: %s", sorted(unknown) + ) + # Dedup keys to minimise per-task anti-join hash table size. keys_table = keys_table.group_by(upsert_cols).aggregate([]) @@ -371,7 +405,10 @@ def _commit_upsert_scan_merge( # plan_files() reads only manifest metadata — no Parquet data on the driver. t0 = time.perf_counter() - file_scan_tasks = list(self._table.scan(row_filter=coarse_filter).plan_files()) + scan = self._table.scan(row_filter=coarse_filter, case_sensitive=case_sensitive) + scan = scan.use_ref(branch) + file_scan_tasks = list(scan.plan_files()) + logger.info( "[scan-merge] planned %d candidate file(s) in %.2fs", len(file_scan_tasks), @@ -379,8 +416,8 @@ def _commit_upsert_scan_merge( ) if not file_scan_tasks: - # No existing files match the coarse filter — pure inserts only. - self._append_and_commit(txn, data_files) + # No existing files match the coarse filter, so it's a pure insert. + self._append_and_commit(txn, data_files, branch=branch) return # Put the deduped keys in the object store once; all tasks share one copy. @@ -433,11 +470,11 @@ def _commit_upsert_scan_merge( n_untouched, ) - # Single atomic commit: schema update (already staged in txn) + this overwrite. + # Single atomic commit: schema update (already staged in txn), and overwrite. # _OverwriteFiles handles both file-level deletes and appends in one snapshot. t0 = time.perf_counter() with txn.update_snapshot( - snapshot_properties=self._snapshot_properties + snapshot_properties=self._snapshot_properties, branch=branch ).overwrite() as snap: for old_file, fp_files in results: if old_file is not None: @@ -454,15 +491,22 @@ def _commit_upsert_scan_merge( logger.info("[scan-merge] committed in %.2fs", time.perf_counter() - t0) def _append_and_commit( - self, txn: "Table.transaction", data_files: List["DataFile"] + self, + txn: "Table.transaction", + data_files: List["DataFile"], + branch: str = "main", ) -> None: """Append data files to a transaction and commit. Args: txn: PyIceberg transaction object data_files: List of DataFile objects to append + branch: Iceberg branch to commit the snapshot to. Defaults to "main" + to match pyiceberg's default """ - with txn._append_snapshot_producer(self._snapshot_properties) as append_files: + with txn._append_snapshot_producer( + self._snapshot_properties, branch=branch + ) as append_files: for data_file in data_files: append_files.append_data_file(data_file) @@ -514,11 +558,7 @@ def _commit_upsert( # Only delete if we have non-NULL keys if len(keys_table) > 0: - delete_kwargs = self._upsert_kwargs.copy() - delete_kwargs.pop(_UPSERT_COLS_ID, None) - self._commit_upsert_scan_merge( - txn, data_files, keys_table, upsert_cols, delete_kwargs - ) + self._commit_upsert_scan_merge(txn, data_files, keys_table, upsert_cols) return else: logger.info("[upsert commit] No upsert keys — skipping delete phase") @@ -529,7 +569,8 @@ def _commit_upsert( len(data_files), ) t0 = time.perf_counter() - self._append_and_commit(txn, data_files) + branch = self._upsert_kwargs.get("branch", "main") + self._append_and_commit(txn, data_files, branch=branch) logger.info( "[upsert commit] Append+commit done in %.2fs", time.perf_counter() - t0, From b81fe8836cd2eee694f6c5c004d12486f5637cc8 Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Tue, 19 May 2026 14:51:55 -0700 Subject: [PATCH 06/14] [Data] Add new tests for upserts - case insensitive arg, and correctness on new join column Signed-off-by: Ayush Kumar --- .../ray/data/tests/datasource/test_iceberg.py | 103 ++++++++++++++++++ 1 file changed, 103 insertions(+) diff --git a/python/ray/data/tests/datasource/test_iceberg.py b/python/ray/data/tests/datasource/test_iceberg.py index 5f126b82848f..770e8d69bb81 100644 --- a/python/ray/data/tests/datasource/test_iceberg.py +++ b/python/ray/data/tests/datasource/test_iceberg.py @@ -1559,6 +1559,109 @@ def test_upsert_string_key(self, clean_table): ) assert rows_same(result, expected) + def test_upsert_case_insensitive_join_cols(self, clean_table): + """``case_sensitive=False`` should let join_cols match table columns + whose casing differs from the supplied names.""" + from ray.data import SaveMode + + seed = _create_typed_dataframe( + { + "col_a": [1, 2, 3, 4, 5], + "col_b": ["seed_1", "seed_2", "seed_3", "seed_4", "seed_5"], + "col_c": [1] * 5, + } + ) + _write_to_iceberg(seed) + + upsert_data = _create_typed_dataframe( + { + "col_a": [1, 5, 6], + "col_b": ["updated_1", "updated_5", "new_6"], + "col_c": [1, 1, 1], + } + ) + _write_to_iceberg( + upsert_data, + mode=SaveMode.UPSERT, + upsert_kwargs={"join_cols": ["COL_A"], "case_sensitive": False}, + ) + + result = _read_from_iceberg(sort_by="col_a") + expected = _create_typed_dataframe( + { + "col_a": [1, 2, 3, 4, 5, 6], + "col_b": [ + "updated_1", + "seed_2", + "seed_3", + "seed_4", + "updated_5", + "new_6", + ], + "col_c": [1] * 6, + } + ) + assert rows_same(result, expected) + + def test_upsert_with_new_column(self, clean_table): + """Upsert that introduces a new column must evolve the table schema, + populate the new column for upserted rows, and leave NULLs for + untouched seed rows (including false-positive rows preserved during + rewrite).""" + from ray.data import SaveMode + + seed = _create_typed_dataframe( + { + "col_a": list(range(1, 6)), + "col_b": [f"seed_{i}" for i in range(1, 6)], + "col_c": [1] * 5, + } + ) + _write_to_iceberg(seed) + + # Upsert touches col_a=1 and col_a=5 (false positives at 2, 3, 4 in + # the same file) and introduces a new column ``col_d``. + upsert_data = _create_typed_dataframe( + { + "col_a": [1, 5, 6], + "col_b": ["updated_1", "updated_5", "new_6"], + "col_c": [1, 1, 1], + "col_d": ["d_1", "d_5", "d_6"], + } + ) + _write_to_iceberg( + upsert_data, + mode=SaveMode.UPSERT, + upsert_kwargs={"join_cols": ["col_a"]}, + ) + + _verify_schema( + { + "col_a": pyi_types.IntegerType, + "col_b": pyi_types.StringType, + "col_c": pyi_types.IntegerType, + "col_d": pyi_types.StringType, + } + ) + + result = _read_from_iceberg(sort_by="col_a") + expected = _create_typed_dataframe( + { + "col_a": [1, 2, 3, 4, 5, 6], + "col_b": [ + "updated_1", + "seed_2", + "seed_3", + "seed_4", + "updated_5", + "new_6", + ], + "col_c": [1] * 6, + "col_d": ["d_1", None, None, None, "d_5", "d_6"], + } + ) + assert rows_same(result, expected) + @pytest.fixture def table_with_identifier_fields() -> Generator[Tuple[Catalog, Table], None, None]: From 7e514f40bec69df939a7b6c033979c07aff5d24a Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Tue, 19 May 2026 16:58:08 -0700 Subject: [PATCH 07/14] [Data] Add typing and use defined constant for parquet expansion factor in iceberg upsert Signed-off-by: Ayush Kumar --- .../_internal/datasource/iceberg_datasink.py | 18 ++++++++---------- 1 file changed, 8 insertions(+), 10 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index 999f2a363464..2fa4b0a865c3 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -14,6 +14,7 @@ from ray.data.datasource.datasink import Datasink, WriteResult from ray.data.expressions import Expr from ray.util.annotations import DeveloperAPI +from ray.data._internal.datasource.parquet_datasource import PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT if TYPE_CHECKING: import pyarrow as pa @@ -22,7 +23,7 @@ from pyiceberg.io import FileIO from pyiceberg.manifest import DataFile from pyiceberg.schema import Schema - from pyiceberg.table import FileScanTask, Table + from pyiceberg.table import DataScan, FileScanTask, Table from pyiceberg.table.metadata import TableMetadata from pyiceberg.table.update.schema import UpdateSchema @@ -403,11 +404,13 @@ def _commit_upsert_scan_merge( coarse_filter = self._build_coarse_range_filter(keys_table, upsert_cols) logger.debug("[scan-merge] coarse_filter=%s", coarse_filter) - # plan_files() reads only manifest metadata — no Parquet data on the driver. + # plan_files() reads only manifest metadata, no Parquet data on the driver. t0 = time.perf_counter() - scan = self._table.scan(row_filter=coarse_filter, case_sensitive=case_sensitive) + scan: "DataScan" = self._table.scan( + row_filter=coarse_filter, case_sensitive=case_sensitive + ) scan = scan.use_ref(branch) - file_scan_tasks = list(scan.plan_files()) + file_scan_tasks: List["FileScanTask"] = list(scan.plan_files()) logger.info( "[scan-merge] planned %d candidate file(s) in %.2fs", @@ -423,15 +426,10 @@ def _commit_upsert_scan_merge( # Put the deduped keys in the object store once; all tasks share one copy. keys_ref = ray.put(keys_table) - # Parquet decompression factor: compressed Parquet typically expands - # ~8-10x in memory once decoded. Conservative default keeps Ray's - # scheduler from stacking too many rewrite tasks on one node. - _PARQUET_EXPANSION = 8 - t0 = time.perf_counter() refs = [ _rewrite_iceberg_file.options( - memory=int(task.file.file_size_in_bytes * _PARQUET_EXPANSION) + memory=int(task.file.file_size_in_bytes * PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT) ).remote(task, keys_ref, upsert_cols, self._table_metadata, self._io) for task in file_scan_tasks ] From 9cd851c8e86ef223a06d6855b4007c3078b691cb Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Tue, 19 May 2026 17:52:33 -0700 Subject: [PATCH 08/14] [Data] Added stall timeout for rewrite tasks in iceberg upsert and renamed fp_ to preserved_ Signed-off-by: Ayush Kumar --- .../_internal/datasource/iceberg_datasink.py | 72 +++++++++++-------- .../ray/data/tests/datasource/test_iceberg.py | 21 +++--- 2 files changed, 54 insertions(+), 39 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index 2fa4b0a865c3..4a93d372a4ab 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -7,6 +7,9 @@ import ray from ray._common.retry import call_with_retry +from ray.data._internal.datasource.parquet_datasource import ( + PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT, +) from ray.data._internal.execution.interfaces import TaskContext from ray.data._internal.savemode import SaveMode from ray.data.block import Block, BlockAccessor @@ -14,7 +17,6 @@ from ray.data.datasource.datasink import Datasink, WriteResult from ray.data.expressions import Expr from ray.util.annotations import DeveloperAPI -from ray.data._internal.datasource.parquet_datasource import PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT if TYPE_CHECKING: import pyarrow as pa @@ -29,6 +31,7 @@ logger = logging.getLogger(__name__) +_REWRITE_STALL_TIMEOUT_S = 600 @ray.remote def _rewrite_iceberg_file( @@ -38,14 +41,14 @@ def _rewrite_iceberg_file( table_metadata: "TableMetadata", io: "FileIO", ) -> "tuple[Optional[DataFile], List[DataFile]]": - """Read one Iceberg file, anti-join against upsert keys, write false-positive rows. + """Read one Iceberg file, anti-join against upsert keys, write preserved rows. - False positives are rows in the file that are not in the upsert batch. The + Preserved rows are rows in the file that are not in the upsert batch. The coarse range filter would delete them (see ``IcebergDatasink._build_coarse_range_filter``), so we preserve them by writing them as new data files before the delete. - Returns (original DataFile to delete, list of new FP DataFiles). - If the entire file is matched (no FPs), returns (file, []). + Returns (original DataFile to delete, list of new preserved DataFiles). + If the entire file is matched (no preserved rows), returns (file, []). If the file has no matched rows at all, returns (None, []), leave it untouched. """ import hashlib @@ -91,12 +94,12 @@ def _rewrite_iceberg_file( ) idx_col = pa.array(range(len(batch)), type=pa.int64()) - fp_keys = batch_keys.append_column("__row_idx__", idx_col).join( + preserved_keys = batch_keys.append_column("__row_idx__", idx_col).join( keys_ref, keys=upsert_cols, join_type="left anti" ) - if len(fp_keys) == 0: - # Every row in this file is being upserted — delete the whole file, no FP file needed. + if len(preserved_keys) == 0: + # Every row in this file is being upserted — delete the whole file, no preserved file needed. logger.debug( "[rewrite] %s: all %d rows matched -> whole-file delete", file_path.split("/")[-1], @@ -104,31 +107,34 @@ def _rewrite_iceberg_file( ) return (file_scan_task.file, []) - if len(fp_keys) == len(batch): + if len(preserved_keys) == len(batch): # No rows in this file match any upsert key — leave it alone entirely. logger.debug( "[rewrite] %s: 0 rows matched -> untouched", file_path.split("/")[-1] ) return (None, []) - fp_rows = batch.take(fp_keys["__row_idx__"]) + preserved_rows = batch.take(preserved_keys["__row_idx__"]) # Derive a deterministic write_uuid from the source file path so that # task retries overwrite the same object rather than leaking orphan files. - fp_write_uuid = _uuid.UUID(hashlib.md5(file_path.encode()).hexdigest()) - fp_files = list( + preserved_write_uuid = _uuid.UUID(hashlib.md5(file_path.encode()).hexdigest()) + preserved_files = list( _dataframe_to_data_files( - table_metadata=table_metadata, df=fp_rows, io=io, write_uuid=fp_write_uuid + table_metadata=table_metadata, + df=preserved_rows, + io=io, + write_uuid=preserved_write_uuid, ) ) logger.debug( - "[rewrite] %s: %d/%d rows are FPs -> wrote %d FP file(s) in %.2fs", + "[rewrite] %s: %d/%d rows preserved -> wrote %d preserved file(s) in %.2fs", file_path.split("/")[-1], - len(fp_keys), + len(preserved_keys), len(batch), - len(fp_files), + len(preserved_files), _time.perf_counter() - t_read, ) - return (file_scan_task.file, fp_files) + return (file_scan_task.file, preserved_files) @dataclass @@ -340,7 +346,7 @@ def _build_coarse_range_filter( """Build an O(1) coarse range filter covering all upsert key values. For each upsert column computes AND(GTE(col, min), LTE(col, max)). - The filter may match rows outside the upsert batch (false positives); + The filter may match rows outside the upsert batch (filter overshoot); callers must anti-join to identify and preserve those rows. """ import pyarrow.compute as pc @@ -378,11 +384,11 @@ def _commit_upsert_scan_merge( 1. Build an O(1) coarse range filter using min-max covering upsert key values (for each column). 2. plan_files() on the driver to find candidate files that could be updated 3. Dispatch one Ray task per candidate file. Each task reads its file, - anti-joins against the upsert keys to find false positives (rows that + anti-joins against the upsert keys to find preserved rows (rows that the coarse delete would remove but that are NOT being upserted), and writes them as new data files directly to storage. 4. Commit atomically via txn.update_snapshot().overwrite(): delete each - original candidate file and append FP files + new upsert data files. + original candidate file and append preserved files + new upsert data files. """ import time @@ -409,6 +415,7 @@ def _commit_upsert_scan_merge( scan: "DataScan" = self._table.scan( row_filter=coarse_filter, case_sensitive=case_sensitive ) + # Use the specific branch for the scan scan = scan.use_ref(branch) file_scan_tasks: List["FileScanTask"] = list(scan.plan_files()) @@ -429,7 +436,10 @@ def _commit_upsert_scan_merge( t0 = time.perf_counter() refs = [ _rewrite_iceberg_file.options( - memory=int(task.file.file_size_in_bytes * PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT) + memory=int( + task.file.file_size_in_bytes + * PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT + ) ).remote(task, keys_ref, upsert_cols, self._table_metadata, self._io) for task in file_scan_tasks ] @@ -441,7 +451,8 @@ def _commit_upsert_scan_merge( _LOG_INTERVAL = max(1, len(refs) // 10) # log ~10 times total while pending: done, pending = ray.wait( - pending, num_returns=min(_LOG_INTERVAL, len(pending)) + pending, num_returns=min(_LOG_INTERVAL, len(pending)), + timeout=_REWRITE_STALL_TIMEOUT_S, fetch_local=True ) results.extend(ray.get(done)) logger.debug( @@ -458,9 +469,14 @@ def _commit_upsert_scan_merge( ) # Count how many files were wholly deleted vs partially rewritten. - n_whole_delete = sum(1 for old, fps in results if old is not None and not fps) - n_partial = sum(1 for old, fps in results if fps) - n_untouched = sum(1 for old, fps in results if old is None) + n_whole_delete = n_partial = n_untouched = 0 + for old, preserved_files in results: + if old is None: + n_untouched += 1 + elif preserved_files: + n_partial += 1 + else: + n_whole_delete += 1 logger.info( "[scan-merge] files: %d whole-delete, %d partial-rewrite, %d untouched", n_whole_delete, @@ -474,11 +490,11 @@ def _commit_upsert_scan_merge( with txn.update_snapshot( snapshot_properties=self._snapshot_properties, branch=branch ).overwrite() as snap: - for old_file, fp_files in results: + for old_file, preserved_files in results: if old_file is not None: snap.delete_data_file(old_file) - for fp_file in fp_files: - snap.append_data_file(fp_file) + for preserved_file in preserved_files: + snap.append_data_file(preserved_file) for df in data_files: snap.append_data_file(df) diff --git a/python/ray/data/tests/datasource/test_iceberg.py b/python/ray/data/tests/datasource/test_iceberg.py index 770e8d69bb81..5eaf6fa85dfb 100644 --- a/python/ray/data/tests/datasource/test_iceberg.py +++ b/python/ray/data/tests/datasource/test_iceberg.py @@ -1311,9 +1311,9 @@ class TestUpsertScanMerge: See ``IcebergDatasink._commit_upsert_scan_merge`` for algorithm details. """ - def test_upsert_preserves_false_positives_sparse_keys(self, clean_table): - """Sparse upsert keys leave intermediate rows as false positives that - must be preserved after the rewrite.""" + def test_upsert_preserves_rows_sparse_keys(self, clean_table): + """Sparse upsert keys leave intermediate rows that must be preserved + after the rewrite.""" from ray.data import SaveMode seed = _create_typed_dataframe( @@ -1352,7 +1352,7 @@ def test_upsert_preserves_false_positives_sparse_keys(self, clean_table): def test_upsert_across_multiple_files(self, clean_table): """Two separate seed writes produce at least two data files. A sparse - upsert that spans both files must rewrite false positives in each.""" + upsert that spans both files must preserve non-upsert rows in each.""" from ray.data import SaveMode _write_to_iceberg( @@ -1479,10 +1479,10 @@ def test_upsert_pure_insert_short_circuit(self, clean_table): ) assert rows_same(result, expected) - def test_upsert_composite_key_preserves_false_positives(self, clean_table): + def test_upsert_composite_key_preserves_rows(self, clean_table): """Composite-key anti-join must match on all join columns; rows that - share one column with an upsert key but not the full composite are - false positives that must be preserved.""" + share one column with an upsert key but not the full composite must be + preserved.""" from ray.data import SaveMode composites = [(a, b) for a in [1, 2, 3] for b in ["x", "y", "z"]] @@ -1606,8 +1606,7 @@ def test_upsert_case_insensitive_join_cols(self, clean_table): def test_upsert_with_new_column(self, clean_table): """Upsert that introduces a new column must evolve the table schema, populate the new column for upserted rows, and leave NULLs for - untouched seed rows (including false-positive rows preserved during - rewrite).""" + untouched seed rows (including preserved rows rewritten during upsert).""" from ray.data import SaveMode seed = _create_typed_dataframe( @@ -1619,8 +1618,8 @@ def test_upsert_with_new_column(self, clean_table): ) _write_to_iceberg(seed) - # Upsert touches col_a=1 and col_a=5 (false positives at 2, 3, 4 in - # the same file) and introduces a new column ``col_d``. + # Upsert touches col_a=1 and col_a=5 (preserved rows at 2, 3, 4 in the + # same file) and introduces a new column ``col_d``. upsert_data = _create_typed_dataframe( { "col_a": [1, 5, 6], From bec146d7c2781f12d57df6d156cfebad030c6472 Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Tue, 19 May 2026 18:01:12 -0700 Subject: [PATCH 09/14] [Data] Add diagram to explain task based iceberg upsert Signed-off-by: Ayush Kumar --- .../_internal/datasource/iceberg_datasink.py | 30 +++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index 4a93d372a4ab..503653137397 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -381,6 +381,36 @@ def _commit_upsert_scan_merge( ) -> None: """Upsert commit using coarse range filter + per-file distributed anti-join. + ┌─────────────────────────────────────────────────────────────┐ + │ Stage 1: Build coarse filter (driver) │ + │ keys_table ──► min/max per col ──► coarse_filter │ + └─────────────────────────────────────────────────────────────┘ + │ + ▼ + ┌─────────────────────────────────────────────────────────────┐ + │ Stage 2: Plan candidate files (driver) │ + │ table.scan(coarse_filter).plan_files() │ + │ ──► file_scan_tasks │ + └─────────────────────────────────────────────────────────────┘ + │ + ▼ + ┌─────────────────────────────────────────────────────────────┐ + │ Stage 3: Rewrite (one _rewrite_iceberg_file task per file) │ + │ read file ─► anti-join keys ─► write preserved rows │ + │ returns (old_file, preserved_files) │ + └─────────────────────────────────────────────────────────────┘ + │ + ▼ + ┌─────────────────────────────────────────────────────────────┐ + │ Stage 4: Atomic overwrite (driver) │ + │ delete old_file (each rewritten candidate) │ + │ append preserved_files (preserved rows kept) │ + │ append data_files (new upsert payload) │ + └─────────────────────────────────────────────────────────────┘ + │ + ▼ + commit_transaction + 1. Build an O(1) coarse range filter using min-max covering upsert key values (for each column). 2. plan_files() on the driver to find candidate files that could be updated 3. Dispatch one Ray task per candidate file. Each task reads its file, From ddba5b8ff300323b23ad79f988a82df4d6ce3303 Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Wed, 20 May 2026 13:58:14 -0700 Subject: [PATCH 10/14] [Data] Make rewrite join streaming for iceberg upsert Signed-off-by: Ayush Kumar --- .../_internal/datasource/iceberg_datasink.py | 85 +++++++++++++------ 1 file changed, 57 insertions(+), 28 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index 503653137397..fa687d721a1e 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -47,6 +47,11 @@ def _rewrite_iceberg_file( coarse range filter would delete them (see ``IcebergDatasink._build_coarse_range_filter``), so we preserve them by writing them as new data files before the delete. + The file is read in streaming fashion via ``ArrowScan.to_record_batches()`` + so the full file is never materialised at once. The anti-join is applied + per RecordBatch and preserved rows are accumulated, then concatenated and + written as a single output once the stream is exhausted. + Returns (original DataFile to delete, list of new preserved DataFiles). If the entire file is matched (no preserved rows), returns (file, []). If the file has no matched rows at all, returns (None, []), leave it untouched. @@ -63,58 +68,81 @@ def _rewrite_iceberg_file( file_size_mb = file_scan_task.file.file_size_in_bytes / 1e6 t_start = _time.perf_counter() - batch = ArrowScan( + # Cast targets pulled from keys_ref once — applied per batch so PyArrow's join + # doesn't raise ArrowInvalid on utf8/large_utf8 or similar width mismatches. + key_cast = {f.name: f.type for f in keys_ref.schema if f.name in upsert_cols} + + record_batches = ArrowScan( table_metadata=table_metadata, io=io, projected_schema=table_metadata.schema(), row_filter=AlwaysTrue(), - ).to_table(tasks=[file_scan_task]) + ).to_record_batches(tasks=[file_scan_task]) + + preserved_batches: List["pa.Table"] = [] + total_in_rows = 0 + total_preserved_rows = 0 + n_batches = 0 + + for rb in record_batches: + n_batches += 1 + batch_table = pa.Table.from_batches([rb]) + if len(batch_table) == 0: + continue + total_in_rows += len(batch_table) + + batch_keys = batch_table.select(upsert_cols) + for col_name, target_type in key_cast.items(): + if batch_keys.schema.field(col_name).type != target_type: + col_idx = batch_keys.schema.get_field_index(col_name) + batch_keys = batch_keys.set_column( + col_idx, col_name, batch_keys[col_name].cast(target_type) + ) + + idx_col = pa.array(range(len(batch_table)), type=pa.int64()) + preserved_keys = batch_keys.append_column("__row_idx__", idx_col).join( + keys_ref, keys=upsert_cols, join_type="left anti" + ) + + if len(preserved_keys) > 0: + preserved_batches.append(batch_table.take(preserved_keys["__row_idx__"])) + total_preserved_rows += len(preserved_keys) t_read = _time.perf_counter() logger.debug( - "[rewrite] read %d rows / %.1f MB (compressed) from %s in %.2fs", - len(batch), + "[rewrite] stream-read+join %d rows / %.1f MB (compressed) from %s " + "across %d batch(es) in %.2fs", + total_in_rows, file_size_mb, file_path.split("/")[-1], + n_batches, t_read - t_start, ) - if len(batch) == 0: + if total_in_rows == 0: return (None, []) - # Cast batch key columns to match keys_ref types so PyArrow's join doesn't - # raise ArrowInvalid on utf8/large_utf8 or similar width mismatches. - key_cast = {f.name: f.type for f in keys_ref.schema if f.name in upsert_cols} - batch_keys = batch.select(upsert_cols) - for col_name, target_type in key_cast.items(): - if batch_keys.schema.field(col_name).type != target_type: - col_idx = batch_keys.schema.get_field_index(col_name) - batch_keys = batch_keys.set_column( - col_idx, col_name, batch_keys[col_name].cast(target_type) - ) - - idx_col = pa.array(range(len(batch)), type=pa.int64()) - preserved_keys = batch_keys.append_column("__row_idx__", idx_col).join( - keys_ref, keys=upsert_cols, join_type="left anti" - ) - - if len(preserved_keys) == 0: + if total_preserved_rows == 0: # Every row in this file is being upserted — delete the whole file, no preserved file needed. logger.debug( "[rewrite] %s: all %d rows matched -> whole-file delete", file_path.split("/")[-1], - len(batch), + total_in_rows, ) return (file_scan_task.file, []) - if len(preserved_keys) == len(batch): + if total_preserved_rows == total_in_rows: # No rows in this file match any upsert key — leave it alone entirely. logger.debug( "[rewrite] %s: 0 rows matched -> untouched", file_path.split("/")[-1] ) return (None, []) - preserved_rows = batch.take(preserved_keys["__row_idx__"]) + # promote_options="permissive" mirrors pyiceberg's own to_table() concat + # (pyiceberg/io/pyarrow.py) to tolerate per-batch schema drift (e.g. utf8 vs large_utf8). + preserved_rows = pa.concat_tables( + preserved_batches, promote_options="permissive" + ) # Derive a deterministic write_uuid from the source file path so that # task retries overwrite the same object rather than leaking orphan files. preserved_write_uuid = _uuid.UUID(hashlib.md5(file_path.encode()).hexdigest()) @@ -129,8 +157,8 @@ def _rewrite_iceberg_file( logger.debug( "[rewrite] %s: %d/%d rows preserved -> wrote %d preserved file(s) in %.2fs", file_path.split("/")[-1], - len(preserved_keys), - len(batch), + total_preserved_rows, + total_in_rows, len(preserved_files), _time.perf_counter() - t_read, ) @@ -469,7 +497,8 @@ def _commit_upsert_scan_merge( memory=int( task.file.file_size_in_bytes * PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT - ) + ), + num_cpus=1 ).remote(task, keys_ref, upsert_cols, self._table_metadata, self._io) for task in file_scan_tasks ] From d3286fe813ba73e0ab8f2fcb0a83a97c1f7b32a6 Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Thu, 21 May 2026 11:18:57 -0700 Subject: [PATCH 11/14] [Data] Bump down upsert rewrite memory hint to file size Signed-off-by: Ayush Kumar --- python/ray/data/_internal/datasource/iceberg_datasink.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index fa687d721a1e..6fac7c940bda 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -7,9 +7,6 @@ import ray from ray._common.retry import call_with_retry -from ray.data._internal.datasource.parquet_datasource import ( - PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT, -) from ray.data._internal.execution.interfaces import TaskContext from ray.data._internal.savemode import SaveMode from ray.data.block import Block, BlockAccessor @@ -496,7 +493,6 @@ def _commit_upsert_scan_merge( _rewrite_iceberg_file.options( memory=int( task.file.file_size_in_bytes - * PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT ), num_cpus=1 ).remote(task, keys_ref, upsert_cols, self._table_metadata, self._io) From d1eb7545fdf748d3de783640f947597eed2c2f70 Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Thu, 21 May 2026 14:02:44 -0700 Subject: [PATCH 12/14] [Data] Bump up the resource min for upsert rewrite tasks to maintain scheduling guarantee Signed-off-by: Ayush Kumar --- python/ray/data/_internal/datasource/iceberg_datasink.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index 6fac7c940bda..fa687d721a1e 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -7,6 +7,9 @@ import ray from ray._common.retry import call_with_retry +from ray.data._internal.datasource.parquet_datasource import ( + PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT, +) from ray.data._internal.execution.interfaces import TaskContext from ray.data._internal.savemode import SaveMode from ray.data.block import Block, BlockAccessor @@ -493,6 +496,7 @@ def _commit_upsert_scan_merge( _rewrite_iceberg_file.options( memory=int( task.file.file_size_in_bytes + * PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT ), num_cpus=1 ).remote(task, keys_ref, upsert_cols, self._table_metadata, self._io) From 70c8ed081d15b8e0114c75611af02ea0f3240021 Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Fri, 22 May 2026 15:53:54 -0700 Subject: [PATCH 13/14] [Data] Cast in iceberg upsert rewrite using batch.cast(), use MiB Signed-off-by: Ayush Kumar --- .../_internal/datasource/iceberg_datasink.py | 41 +++++++++---------- 1 file changed, 19 insertions(+), 22 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index fa687d721a1e..98d999aa0359 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -12,6 +12,7 @@ ) from ray.data._internal.execution.interfaces import TaskContext from ray.data._internal.savemode import SaveMode +from ray.data._internal.util import MiB from ray.data.block import Block, BlockAccessor from ray.data.context import DataContext from ray.data.datasource.datasink import Datasink, WriteResult @@ -33,6 +34,7 @@ _REWRITE_STALL_TIMEOUT_S = 600 + @ray.remote def _rewrite_iceberg_file( file_scan_task: "FileScanTask", @@ -65,12 +67,15 @@ def _rewrite_iceberg_file( from pyiceberg.io.pyarrow import ArrowScan, _dataframe_to_data_files file_path = file_scan_task.file.file_path - file_size_mb = file_scan_task.file.file_size_in_bytes / 1e6 + file_name = file_path.split("/")[-1] + file_size_mb = file_scan_task.file.file_size_in_bytes / MiB t_start = _time.perf_counter() - # Cast targets pulled from keys_ref once — applied per batch so PyArrow's join + # Cast target pulled from keys_ref once. Applied per batch so PyArrow's join # doesn't raise ArrowInvalid on utf8/large_utf8 or similar width mismatches. - key_cast = {f.name: f.type for f in keys_ref.schema if f.name in upsert_cols} + target_key_schema = pa.schema( + [keys_ref.schema.field(c) for c in upsert_cols] + ) record_batches = ArrowScan( table_metadata=table_metadata, @@ -91,13 +96,7 @@ def _rewrite_iceberg_file( continue total_in_rows += len(batch_table) - batch_keys = batch_table.select(upsert_cols) - for col_name, target_type in key_cast.items(): - if batch_keys.schema.field(col_name).type != target_type: - col_idx = batch_keys.schema.get_field_index(col_name) - batch_keys = batch_keys.set_column( - col_idx, col_name, batch_keys[col_name].cast(target_type) - ) + batch_keys = batch_table.select(upsert_cols).cast(target_key_schema) idx_col = pa.array(range(len(batch_table)), type=pa.int64()) preserved_keys = batch_keys.append_column("__row_idx__", idx_col).join( @@ -114,7 +113,7 @@ def _rewrite_iceberg_file( "across %d batch(es) in %.2fs", total_in_rows, file_size_mb, - file_path.split("/")[-1], + file_name, n_batches, t_read - t_start, ) @@ -126,23 +125,19 @@ def _rewrite_iceberg_file( # Every row in this file is being upserted — delete the whole file, no preserved file needed. logger.debug( "[rewrite] %s: all %d rows matched -> whole-file delete", - file_path.split("/")[-1], + file_name, total_in_rows, ) return (file_scan_task.file, []) if total_preserved_rows == total_in_rows: # No rows in this file match any upsert key — leave it alone entirely. - logger.debug( - "[rewrite] %s: 0 rows matched -> untouched", file_path.split("/")[-1] - ) + logger.debug("[rewrite] %s: 0 rows matched -> untouched", file_name) return (None, []) # promote_options="permissive" mirrors pyiceberg's own to_table() concat # (pyiceberg/io/pyarrow.py) to tolerate per-batch schema drift (e.g. utf8 vs large_utf8). - preserved_rows = pa.concat_tables( - preserved_batches, promote_options="permissive" - ) + preserved_rows = pa.concat_tables(preserved_batches, promote_options="permissive") # Derive a deterministic write_uuid from the source file path so that # task retries overwrite the same object rather than leaking orphan files. preserved_write_uuid = _uuid.UUID(hashlib.md5(file_path.encode()).hexdigest()) @@ -156,7 +151,7 @@ def _rewrite_iceberg_file( ) logger.debug( "[rewrite] %s: %d/%d rows preserved -> wrote %d preserved file(s) in %.2fs", - file_path.split("/")[-1], + file_name, total_preserved_rows, total_in_rows, len(preserved_files), @@ -498,7 +493,7 @@ def _commit_upsert_scan_merge( task.file.file_size_in_bytes * PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT ), - num_cpus=1 + num_cpus=1, ).remote(task, keys_ref, upsert_cols, self._table_metadata, self._io) for task in file_scan_tasks ] @@ -510,8 +505,10 @@ def _commit_upsert_scan_merge( _LOG_INTERVAL = max(1, len(refs) // 10) # log ~10 times total while pending: done, pending = ray.wait( - pending, num_returns=min(_LOG_INTERVAL, len(pending)), - timeout=_REWRITE_STALL_TIMEOUT_S, fetch_local=True + pending, + num_returns=min(_LOG_INTERVAL, len(pending)), + timeout=_REWRITE_STALL_TIMEOUT_S, + fetch_local=True, ) results.extend(ray.get(done)) logger.debug( From a138c9afdefaae03353d20b70d07c4d35ec6bc97 Mon Sep 17 00:00:00 2001 From: Ayush Kumar Date: Fri, 22 May 2026 17:21:23 -0700 Subject: [PATCH 14/14] [Data] Bump up memory estimate due to materialization in record_batches, refactor iceberg upsert to address comments Signed-off-by: Ayush Kumar --- .../_internal/datasource/iceberg_datasink.py | 21 +++++++++++-------- 1 file changed, 12 insertions(+), 9 deletions(-) diff --git a/python/ray/data/_internal/datasource/iceberg_datasink.py b/python/ray/data/_internal/datasource/iceberg_datasink.py index 98d999aa0359..9216d11d8d0f 100644 --- a/python/ray/data/_internal/datasource/iceberg_datasink.py +++ b/python/ray/data/_internal/datasource/iceberg_datasink.py @@ -62,6 +62,7 @@ def _rewrite_iceberg_file( import time as _time import uuid as _uuid + import numpy as np import pyarrow as pa from pyiceberg.expressions import AlwaysTrue from pyiceberg.io.pyarrow import ArrowScan, _dataframe_to_data_files @@ -73,9 +74,7 @@ def _rewrite_iceberg_file( # Cast target pulled from keys_ref once. Applied per batch so PyArrow's join # doesn't raise ArrowInvalid on utf8/large_utf8 or similar width mismatches. - target_key_schema = pa.schema( - [keys_ref.schema.field(c) for c in upsert_cols] - ) + target_key_schema = pa.schema([keys_ref.schema.field(c) for c in upsert_cols]) record_batches = ArrowScan( table_metadata=table_metadata, @@ -84,7 +83,7 @@ def _rewrite_iceberg_file( row_filter=AlwaysTrue(), ).to_record_batches(tasks=[file_scan_task]) - preserved_batches: List["pa.Table"] = [] + preserved_rows: Optional["pa.Table"] = None total_in_rows = 0 total_preserved_rows = 0 n_batches = 0 @@ -98,13 +97,19 @@ def _rewrite_iceberg_file( batch_keys = batch_table.select(upsert_cols).cast(target_key_schema) - idx_col = pa.array(range(len(batch_table)), type=pa.int64()) + idx_col = pa.array(np.arange(len(batch_table), dtype=np.int64)) preserved_keys = batch_keys.append_column("__row_idx__", idx_col).join( keys_ref, keys=upsert_cols, join_type="left anti" ) if len(preserved_keys) > 0: - preserved_batches.append(batch_table.take(preserved_keys["__row_idx__"])) + new_rows = batch_table.take(preserved_keys["__row_idx__"]) + if preserved_rows is None: + preserved_rows = new_rows + else: + preserved_rows = pa.concat_tables( + [preserved_rows, new_rows], promote_options="permissive" + ) total_preserved_rows += len(preserved_keys) t_read = _time.perf_counter() @@ -135,9 +140,6 @@ def _rewrite_iceberg_file( logger.debug("[rewrite] %s: 0 rows matched -> untouched", file_name) return (None, []) - # promote_options="permissive" mirrors pyiceberg's own to_table() concat - # (pyiceberg/io/pyarrow.py) to tolerate per-batch schema drift (e.g. utf8 vs large_utf8). - preserved_rows = pa.concat_tables(preserved_batches, promote_options="permissive") # Derive a deterministic write_uuid from the source file path so that # task retries overwrite the same object rather than leaking orphan files. preserved_write_uuid = _uuid.UUID(hashlib.md5(file_path.encode()).hexdigest()) @@ -492,6 +494,7 @@ def _commit_upsert_scan_merge( memory=int( task.file.file_size_in_bytes * PARQUET_ENCODING_RATIO_ESTIMATE_DEFAULT + * 3 # Bump memory estimate to account for the anti-join and the preserved rows (also since to_record_batches materializes the entire table in memory, see https://github.com/apache/iceberg-python/issues/3036) ), num_cpus=1, ).remote(task, keys_ref, upsert_cols, self._table_metadata, self._io)