From 22b9b7252df71fb3918d439efeddf30b23c53672 Mon Sep 17 00:00:00 2001 From: fanng <“fanng@apache.org”> Date: Wed, 22 Jul 2026 16:51:41 +0900 Subject: [PATCH] fix: preserve untouched fragments in slow merges --- daft_lance/lance_merge_column.py | 17 +++++---- .../io/lancedb/test_lance_merge_evolution.py | 37 +++++++++++++++++++ 2 files changed, 47 insertions(+), 7 deletions(-) diff --git a/daft_lance/lance_merge_column.py b/daft_lance/lance_merge_column.py index 8984120..c6ea934 100644 --- a/daft_lance/lance_merge_column.py +++ b/daft_lance/lance_merge_column.py @@ -24,6 +24,14 @@ _FRAGMENT_HANDLER_RETURN_DTYPE = DataType.struct({"fragment_meta": DataType.binary(), "schema": DataType.binary()}) +def _include_untouched_fragments(fragment_metas: list[Any], lance_ds: lance.LanceDataset) -> None: + """Add metadata for fragments that were not rewritten by a partial merge.""" + touched_fragment_ids = {int(fragment_meta.to_json()["id"]) for fragment_meta in fragment_metas} + fragment_metas.extend( + fragment.metadata for fragment in lance_ds.get_fragments() if fragment.fragment_id not in touched_fragment_ids + ) + + @daft_cls class FragmentHandler: def __init__( @@ -445,7 +453,6 @@ def _merge_fast_path( commit_messages = grouped.collect().to_pydict()["commit_message"] new_schema = None fragment_metas: list[Any] = [] - enriched_frag_ids: set[int] = set() for commit_message in commit_messages: fragment_meta_bytes = commit_message["fragment_meta"] @@ -454,18 +461,13 @@ def _merge_fast_path( continue fmeta = daft.pickle.loads(fragment_meta_bytes) fragment_metas.append(fmeta) - # pylance 6.0.0 FragmentMetadata exposes the id only via to_json() - enriched_frag_ids.add(int(fmeta.to_json()["id"])) if new_schema is None: new_schema = daft.pickle.loads(schema_bytes) if new_schema is None: raise ValueError("Fast path produced no fragment metadata") - # Include untouched fragments (they'll get NULLs for new columns) - for frag in lance_ds.get_fragments(): - if frag.fragment_id not in enriched_frag_ids: - fragment_metas.append(frag.metadata) + _include_untouched_fragments(fragment_metas, lance_ds) op = lance.LanceOperation.Merge(fragment_metas, LanceSchema.from_pyarrow(new_schema)) return lance.LanceDataset.commit( @@ -515,6 +517,7 @@ def _merge_slow_path( continue if new_schema is None: return lance_ds + _include_untouched_fragments(fragment_metas, lance_ds) op = lance.LanceOperation.Merge(fragment_metas, new_schema) return lance_ds.commit( uri, diff --git a/tests/io/lancedb/test_lance_merge_evolution.py b/tests/io/lancedb/test_lance_merge_evolution.py index 6b9c66a..276ced3 100755 --- a/tests/io/lancedb/test_lance_merge_evolution.py +++ b/tests/io/lancedb/test_lance_merge_evolution.py @@ -1,5 +1,6 @@ from __future__ import annotations +import lance import pytest import daft @@ -140,3 +141,39 @@ def test_merge_columns_df_rowaddr(lance_dataset_path): for a, b in zip(out["lat"], out["double_lat"]): assert pytest.approx(a * 2, rel=1e-6) == b assert df_after.count_rows() == df_loaded.count_rows() + + +def test_merge_columns_df_slow_path_preserves_untouched_fragments(lance_dataset_path): + for fragment_index in range(4): + first_id = fragment_index * 2 + daft.from_pydict( + { + "id": [first_id, first_id + 1], + "value": [first_id * 10, (first_id + 1) * 10], + } + ).write_lance(lance_dataset_path, mode="create" if fragment_index == 0 else "append") + + dataset_before = lance.dataset(lance_dataset_path) + version_before = dataset_before.version + fragment_count_before = len(dataset_before.get_fragments()) + + source = daft.read_lance( + lance_dataset_path, + default_scan_options={"with_row_address": True}, + include_fragment_id=True, + ).where(daft.col("id").is_in([0, 2, 3])) + source = source.with_column("merged_value", daft.col("value") + 1) + + daft_lance.merge_columns_df( + source.select("fragment_id", "_rowaddr", "merged_value"), + lance_dataset_path, + ) + + dataset_after = lance.dataset(lance_dataset_path) + result = dataset_after.to_table().sort_by("id").to_pydict() + assert result["id"] == list(range(8)) + assert result["value"] == [i * 10 for i in range(8)] + assert result["merged_value"] == [1, None, 21, 31, None, None, None, None] + assert dataset_after.count_rows() == 8 + assert len(dataset_after.get_fragments()) == fragment_count_before == 4 + assert dataset_after.version == version_before + 1