diff --git a/README.md b/README.md index 106d6bd..57d5634 100644 --- a/README.md +++ b/README.md @@ -38,6 +38,31 @@ from daft_lance import merge_columns_df merge_columns_df(df, "s3://bucket/my_dataset") ``` +### Conditional Overwrite + +Replace just the rows matching a predicate. One Lance commit deletes them from the existing +table and adds the new data, so readers see either the whole replacement or none of it. + +```python +import daft_lance + +daft_lance.write_lance( + df, + "s3://bucket/events", + mode="insert_overwrite", + overwrite_where="dt = DATE '2026-08-25'", +).collect() +``` + +The table must already exist. `overwrite_where` determines which existing rows are removed; +the input DataFrame is appended as-is. Rows outside `overwrite_where` are not replaced by a +later re-run, so callers should filter the input first when they need idempotent replacement. + +> **Warning:** Lance does not treat a concurrent append or update as conflicting with this +> commit, so rows another writer adds while the overwrite runs survive it even when they match +> `overwrite_where`, without any error. Make sure no other writer touches the table during a +> conditional overwrite. + ### Namespace Tables Address Lance tables through a [Lance Namespace](https://lancedb.github.io/lance-namespace/) diff --git a/daft_lance/_lance.py b/daft_lance/_lance.py index 0ad2231..542e24c 100644 --- a/daft_lance/_lance.py +++ b/daft_lance/_lance.py @@ -632,10 +632,11 @@ def compact_files( def write_lance( df: DataFrame, uri: str | pathlib.Path | None = None, - mode: Literal["create", "append", "overwrite"] = "create", + mode: Literal["create", "append", "overwrite", "insert_overwrite"] = "create", io_config: IOConfig | None = None, schema: Schema | pa.Schema | None = None, *, + overwrite_where: str | None = None, table_id: list[str] | None = None, namespace_impl: str | None = None, namespace_properties: dict[str, str] | None = None, @@ -646,9 +647,16 @@ def write_lance( Args: df: The DataFrame to write. uri: The URI of the Lance table. Mutually exclusive with the namespace parameters. - mode: One of "create", "append", or "overwrite". + mode: One of "create", "append", "overwrite", or "insert_overwrite". + ``"insert_overwrite"`` replaces just the rows matching ``overwrite_where``: one Lance + commit deletes them from the existing table and adds this DataFrame's data, so + readers see either the whole replacement or none of it. It requires an existing + table and is not supported with ``use_mem_wal=True``. io_config: A custom IOConfig to use when accessing Lance data. schema: Desired schema to enforce during write; defaults to the DataFrame schema. + overwrite_where: SQL predicate selecting the rows to replace. Required by, and only valid + with, ``mode="insert_overwrite"``. Uses Lance's SQL filter dialect, e.g. + ``"dt = DATE '2026-08-25'"``. table_id: Table identifier within the namespace, e.g. ["catalog", "schema", "table"]. namespace_impl: Lance Namespace implementation, e.g. "dir" or "rest". namespace_properties: Properties for connecting to the namespace, e.g. @@ -670,6 +678,12 @@ def write_lance( >>> daft_lance.write_lance( ... df, namespace_impl="dir", namespace_properties={"root": "/tmp/tables"}, table_id=["t"] ... ).collect() # doctest: +SKIP + + Replace one day's rows and add this batch in a single commit: + + >>> daft_lance.write_lance( + ... df, "/tmp/events", mode="insert_overwrite", overwrite_where="dt = DATE '2026-08-25'" + ... ).collect() # doctest: +SKIP """ validate_uri_or_namespace(uri, namespace_impl, table_id, namespace_properties) @@ -681,6 +695,7 @@ def write_lance( schema, mode, io_config, + overwrite_where=overwrite_where, table_id=table_id, namespace_impl=namespace_impl, namespace_properties=namespace_properties, diff --git a/daft_lance/lance_data_sink.py b/daft_lance/lance_data_sink.py index 3feef7e..d7ad95c 100644 --- a/daft_lance/lance_data_sink.py +++ b/daft_lance/lance_data_sink.py @@ -43,6 +43,27 @@ logger = logging.getLogger(__name__) +# What the caller asks for, and what the write physically does. ``insert_overwrite`` +# writes exactly like an append -- it only differs at commit time -- so it is +# normalized to "append" once in the constructor. Every mode check outside the +# commit path reads the normalized value, because a check that forgets the new +# mode fails silently (see resolve_storage_version, which would skip the +# storage-version compatibility check entirely). +LancePhysicalWriteMode = Literal["create", "append", "overwrite"] + + +def _dataset_stats(dataset: lance.LanceDataset) -> MicroPartition: + """The single-row write result: dataset stats plus the version just produced.""" + stats = dataset.stats.dataset_stats() + return MicroPartition.from_pydict( + { + "num_fragments": pa.array([stats["num_fragments"]], type=pa.int64()), + "num_deleted_rows": pa.array([stats["num_deleted_rows"]], type=pa.int64()), + "num_small_files": pa.array([stats["num_small_files"]], type=pa.int64()), + "version": pa.array([dataset.version], type=pa.int64()), + } + ) + class LanceDataSink(DataSink[list[FragmentMetadata]]): """WriteSink for writing data to a Lance dataset.""" @@ -51,9 +72,10 @@ def __init__( self, uri: str | pathlib.Path | None, schema: Schema | pa.Schema, - mode: Literal["create", "append", "overwrite"] = "create", + mode: Literal["create", "append", "overwrite", "insert_overwrite"] = "create", io_config: IOConfig | None = None, *, + overwrite_where: str | None = None, table_id: list[str] | None = None, namespace_impl: str | None = None, namespace_properties: dict[str, str] | None = None, @@ -70,11 +92,15 @@ def __init__( ) -> None: self._reject_unsupported_modes(mode, use_legacy_format) self._reject_namespace_mem_wal(namespace_impl, table_id, use_mem_wal) + self._validate_insert_overwrite(mode, overwrite_where, use_mem_wal) validate_uri_or_namespace(uri, namespace_impl, table_id, namespace_properties) if uri is not None and not isinstance(uri, (str, pathlib.Path)): raise TypeError(f"Expected URI to be str or pathlib.Path, got {type(uri)}") self._mode = mode + self._is_insert_overwrite = mode == "insert_overwrite" + self._write_mode: LancePhysicalWriteMode = "append" if mode == "insert_overwrite" else mode + self._overwrite_where = overwrite_where.strip() if overwrite_where is not None else None self._uri = uri self._namespace_impl = namespace_impl self._namespace_properties = namespace_properties @@ -133,17 +159,36 @@ def start(self) -> None: self._data_storage_version = resolve_storage_version( self._requested_storage_version, existing_version, - self._mode, + self._write_mode, ) # Auto-pick up any existing lance.blob.v2 columns when appending so the # write path wraps the matching daft binary columns. - if self._mode == "append" and self._table_schema is not None: + if self._write_mode == "append" and self._table_schema is not None: self._blob.add_columns(detect_blob_v2_columns(self._table_schema)) + if self._is_insert_overwrite: + assert existing is not None, "insert_overwrite requires an existing dataset" + self._validate_overwrite_where_against_table(existing) + # Schema actually written to the dataset (blob columns retyped to lance.blob.v2). self._effective_pyarrow_schema = self._blob.build_effective_schema(self._pyarrow_schema) + def _validate_overwrite_where_against_table(self, dataset: lance.LanceDataset) -> None: + """Fail on the driver, before any data is written, if the filter is unusable. + + Planning a scan is enough to surface parse errors and unknown columns; + without this the write only fails at commit time, after the input has + already been written to storage. + """ + assert self._overwrite_where is not None + try: + dataset.scanner(columns=[], filter=self._overwrite_where, limit=1).explain_plan(True) + except Exception as e: + raise ValueError( + f"overwrite_where={self._overwrite_where!r} is not a valid Lance filter for this table: {e}" + ) from e + @property def _namespace_kwargs(self) -> dict[str, Any]: return get_namespace_kwargs(self._namespace_impl, self._namespace_properties, self._table_id) @@ -164,7 +209,7 @@ def _dataset_uri_arg(self) -> str | None: def _resolve_table(self) -> ResolvedNamespaceTable: if self._uri is not None: return ResolvedNamespaceTable(uri=str(self._uri)) - mode = self._mode if self._mode in ("create", "overwrite") else "read" + mode = self._write_mode if self._write_mode in ("create", "overwrite") else "read" resolved = resolve_namespace_table( namespace_impl=self._namespace_impl, namespace_properties=self._namespace_properties, @@ -188,7 +233,8 @@ def _merged_storage_options(self, resolved: ResolvedNamespaceTable) -> dict[str, @staticmethod def _reject_unsupported_modes( - mode: Literal["create", "append", "overwrite"], use_legacy_format: bool | None + mode: Literal["create", "append", "overwrite", "insert_overwrite"], + use_legacy_format: bool | None, ) -> None: # This mode was never functional and customers must use merge_columns_df. if mode == "merge": # type: ignore[comparison-overlap] @@ -207,6 +253,29 @@ def _reject_unsupported_modes( stacklevel=3, ) + @staticmethod + def _validate_insert_overwrite( + mode: Literal["create", "append", "overwrite", "insert_overwrite"], + overwrite_where: str | None, + use_mem_wal: bool, + ) -> None: + """Conditional overwrite needs a predicate, and only works copy-on-write.""" + if mode != "insert_overwrite": + if overwrite_where is not None: + raise ValueError(f'overwrite_where is only supported with mode="insert_overwrite", got mode="{mode}".') + return + if overwrite_where is None or not overwrite_where.strip(): + raise ValueError( + 'mode="insert_overwrite" requires a non-empty SQL predicate in overwrite_where, ' + "e.g. overwrite_where=\"dt = '2026-08-25'\"." + ) + if use_mem_wal: + raise ValueError( + 'mode="insert_overwrite" is not supported with use_mem_wal=True. The conditional ' + "overwrite commits deletions against a pinned dataset version, which the mem-WAL " + "write path does not go through." + ) + @staticmethod def _reject_namespace_mem_wal(namespace_impl: str | None, table_id: list[str] | None, use_mem_wal: bool) -> None: """Reject the namespace + mem-WAL combination instead of failing mid-write. @@ -274,9 +343,9 @@ def _absorb_existing_dataset(self) -> lance.LanceDataset | None: raise if dataset is None: - if self._mode == "append": - raise ValueError("Cannot append to non-existent Lance dataset.") - if self._mode == "create" and self._storage_options is None and self._table_uri is not None: + if self._write_mode == "append": + raise ValueError(f"Cannot {self._mode} to non-existent Lance dataset.") + if self._write_mode == "create" and self._storage_options is None and self._table_uri is not None: p = pathlib.Path(self._table_uri) if p.is_file(): raise FileExistsError("Target path points to a file, cannot create a dataset here.") @@ -286,13 +355,13 @@ def _absorb_existing_dataset(self) -> lance.LanceDataset | None: self._table_schema = table_schema self._version = dataset.latest_version - if self._mode == "create": + if self._write_mode == "create": raise ValueError( "Cannot create a Lance dataset at a location where one already exists. " 'Use mode="overwrite" to replace it or mode="append" to add to it.' ) - if self._mode == "append" and not _pyarrow_schema_castable( + if self._write_mode == "append" and not _pyarrow_schema_castable( blob_aware_schema_for_validation(self._pyarrow_schema, table_schema), blob_aware_schema_for_validation(table_schema, table_schema), ): @@ -324,7 +393,7 @@ def _write_arrow_table(self, table: pa.Table) -> WriteResult[list[FragmentMetada fragments = lance.fragment.write_fragments( wrapped, dataset_uri=self._table_uri, - mode=self._mode, + mode=self._write_mode, storage_options=self._storage_options, max_rows_per_file=self._max_rows_per_file, max_rows_per_group=self._max_rows_per_group, @@ -429,6 +498,9 @@ def finalize(self, write_results: list[WriteResult[list[FragmentMetadata]]]) -> def _finalize_cow(self, write_results: list[WriteResult[list[FragmentMetadata]]]) -> MicroPartition: fragments = list(chain.from_iterable(write_result.result for write_result in write_results)) + if self._is_insert_overwrite: + return self._finalize_insert_overwrite(fragments) + assert self._effective_pyarrow_schema is not None, "LanceDataSink.start() must run before finalize" operation: lance.LanceOperation.BaseOperation if self._mode == "create" or self._mode == "overwrite": @@ -446,16 +518,47 @@ def _finalize_cow(self, write_results: list[WriteResult[list[FragmentMetadata]]] storage_options=self._storage_options, **self._namespace_commit_kwargs, ) - stats = dataset.stats.dataset_stats() - stats_dict = MicroPartition.from_pydict( - { - "num_fragments": pa.array([stats["num_fragments"]], type=pa.int64()), - "num_deleted_rows": pa.array([stats["num_deleted_rows"]], type=pa.int64()), - "num_small_files": pa.array([stats["num_small_files"]], type=pa.int64()), - "version": pa.array([dataset.version], type=pa.int64()), - } + return _dataset_stats(dataset) + + def _finalize_insert_overwrite(self, fragments: list[FragmentMetadata]) -> MicroPartition: + """Delete the predicate's rows and add this batch's fragments in one commit.""" + from daft_lance.lance_insert_overwrite import apply_insert_overwrite + from daft_lance.namespace import DatasetOpenContext + + assert self._table_uri is not None, "LanceDataSink.start() must run before finalize" + assert self._overwrite_where is not None + + # Pinned to the version start() read: the deletions describe that snapshot, + # and the commit declares it as its read version. + pinned = lance.dataset( + self._dataset_uri_arg, + version=self._version, + storage_options=self._storage_options, + **self._namespace_kwargs, + ) + open_context = DatasetOpenContext.from_dataset( + pinned, + self._table_uri, + storage_options=self._storage_options, + namespace_impl=self._namespace_impl, + namespace_properties=self._namespace_properties, + table_id=self._table_id, + managed_versioning=self._managed_versioning, ) - return stats_dict + dataset = apply_insert_overwrite( + open_context=open_context, + predicate=self._overwrite_where, + new_fragments=fragments, + ) + if dataset is None: + logger.info( + "insert_overwrite matched no rows and wrote no data for predicate %r; no version created", + self._overwrite_where, + ) + dataset = lance.dataset( + self._dataset_uri_arg, storage_options=self._storage_options, **self._namespace_kwargs + ) + return _dataset_stats(dataset) def _finalize_mem_wal(self, write_results: list[WriteResult[list[FragmentMetadata]]]) -> MicroPartition: dataset = lance.dataset(self._dataset_uri_arg, storage_options=self._storage_options, **self._namespace_kwargs) @@ -483,15 +586,7 @@ def _finalize_mem_wal(self, write_results: list[WriteResult[list[FragmentMetadat self._dataset_uri_arg, storage_options=self._storage_options, **self._namespace_kwargs ) - stats = dataset.stats.dataset_stats() - return MicroPartition.from_pydict( - { - "num_fragments": pa.array([stats["num_fragments"]], type=pa.int64()), - "num_deleted_rows": pa.array([stats["num_deleted_rows"]], type=pa.int64()), - "num_small_files": pa.array([stats["num_small_files"]], type=pa.int64()), - "version": pa.array([dataset.version], type=pa.int64()), - } - ) + return _dataset_stats(dataset) class _LanceFragmentBuffer: diff --git a/daft_lance/lance_insert_overwrite.py b/daft_lance/lance_insert_overwrite.py new file mode 100644 index 0000000..010caa1 --- /dev/null +++ b/daft_lance/lance_insert_overwrite.py @@ -0,0 +1,187 @@ +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Any, cast + +import lance +import pyarrow as pa +import pyarrow.compute as pc + +import daft.pickle +from daft import from_pylist +from daft.datatype import DataType +from daft.runners import get_or_create_runner +from daft.udf import cls as daft_cls +from daft.udf import method + +if TYPE_CHECKING: + from lance.fragment import FragmentMetadata + + from daft_lance.namespace import DatasetOpenContext + +logger = logging.getLogger(__name__) + +_FRAGMENT_DELETE_RETURN_DTYPE = DataType.struct( + { + "fragment_id": DataType.int64(), + "fragment_meta": DataType.binary(), + "removed": DataType.bool(), + } +) + +# Pinned to the plan node Lance emits when a scalar index answers the filter. A +# rename would silently cost us fragment pruning, so +# test_pruning_uses_a_scalar_index_when_one_covers_the_predicate asserts on it. +_SCALAR_INDEX_PLAN_MARKER = "ScalarIndexQuery" + +# A Lance row address packs the fragment id into its high 32 bits. +_FRAGMENT_ID_SHIFT = pa.scalar(32, type=pa.uint64()) + +# Each partition builds its own handler and reopens the pinned snapshot once, so +# this bounds the manifest reads a wide table pays for the extra parallelism. +_MAX_DELETE_PARTITIONS = 64 + + +@daft_cls +class FragmentDeleteHandler: + """Applies one delete predicate to a fragment and reports what changed. + + Runs as a Daft UDF: the driver ships fragment ids, each task reopens the + pinned snapshot and writes a deletion file for the rows the predicate + matches. Data files are never rewritten, so row addresses -- and every index + built on them -- stay valid. + """ + + def __init__(self, open_context: DatasetOpenContext, predicate: str) -> None: + self.open_context = open_context + self.predicate = predicate + self._lance_ds: lance.LanceDataset | None = None + + def _dataset(self) -> lance.LanceDataset: + # Opened once per instance, not per fragment: the reopen costs a pinned + # manifest read and must not sit on the per-row path. + if self._lance_ds is None: + self._lance_ds = self.open_context.open_pinned() + return self._lance_ds + + @method.batch(return_dtype=_FRAGMENT_DELETE_RETURN_DTYPE) + def __call__(self, fragment_ids: Any) -> list[dict[str, Any]]: + lance_ds = self._dataset() + results: list[dict[str, Any]] = [] + for fragment_id in fragment_ids: + fragment = lance_ds.get_fragment(fragment_id) + if fragment is None: + raise ValueError(f"Fragment {fragment_id} not found in dataset") + deletions_before = fragment.metadata.num_deletions + updated = fragment.delete(self.predicate) + if updated is None: + # Every row matched: the fragment leaves the dataset entirely. + results.append({"fragment_id": int(fragment_id), "fragment_meta": None, "removed": True}) + continue + # A fragment the predicate missed comes back unchanged; committing it + # as "updated" would only add noise to the transaction. + changed = updated.num_deletions != deletions_before + results.append( + { + "fragment_id": int(fragment_id), + "fragment_meta": daft.pickle.dumps(updated) if changed else None, + "removed": False, + } + ) + return results + + +def _candidate_fragment_ids(dataset: lance.LanceDataset, predicate: str) -> set[int] | None: + """Fragment ids that hold rows matching ``predicate``, or None when unknown. + + Only worth doing when a scalar index can answer the filter: then this is an + index lookup that skips most fragments. Without an index the scan costs the + same full pass the delete step already pays, so we return None and let the + delete visit every fragment rather than paying for both. + """ + scanner = dataset.scanner(columns=[], filter=predicate, with_row_address=True) + if _SCALAR_INDEX_PLAN_MARKER not in scanner.explain_plan(True): + return None + + fragment_ids: set[int] = set() + # Streamed, not to_table(): one overwritten partition can be hundreds of + # millions of row addresses, and we only need the ids they live in. + for batch in scanner.to_batches(): + batch_ids = pc.unique(pc.shift_right(batch.column("_rowaddr"), _FRAGMENT_ID_SHIFT)) + fragment_ids.update(cast("list[int]", batch_ids.to_pylist())) + return fragment_ids + + +def _delete_matching_rows( + open_context: DatasetOpenContext, + predicate: str, + fragment_ids: list[int], +) -> tuple[list[FragmentMetadata], list[int]]: + """Run the per-fragment delete as a Daft job; return (updated, removed).""" + if not fragment_ids: + return [], [] + + df = from_pylist([{"fragment_id": fragment_id} for fragment_id in fragment_ids]) + partitions = min(len(fragment_ids), _MAX_DELETE_PARTITIONS) + # from_pylist lands everything in one partition, which would pin the whole + # delete to a single task on a distributed runner. The native runner has no + # partitions to spread -- repartition there is a no-op that only warns. + if partitions > 1 and get_or_create_runner().name != "native": + df = df.repartition(partitions, "fragment_id") + handler = FragmentDeleteHandler(open_context, predicate) + df = df.with_column("delete_result", handler(df["fragment_id"])) # type: ignore[arg-type] + + updated_fragments: list[FragmentMetadata] = [] + removed_fragment_ids: list[int] = [] + for result in df.collect().to_pydict()["delete_result"]: + if result["removed"]: + removed_fragment_ids.append(int(result["fragment_id"])) + elif result["fragment_meta"] is not None: + updated_fragments.append(daft.pickle.loads(result["fragment_meta"])) + return updated_fragments, removed_fragment_ids + + +def apply_insert_overwrite( + *, + open_context: DatasetOpenContext, + predicate: str, + new_fragments: list[FragmentMetadata], +) -> lance.LanceDataset | None: + """Delete the rows matching ``predicate`` and add ``new_fragments`` in one commit. + + ``open_context`` must be pinned to the version the write started from; that + version is what the commit declares as its read version, so the deletions + describe the snapshot they were computed against. + + Returns the committed dataset, or None when there was nothing to do (no + matching rows and no new data), in which case no version is created. + """ + pinned = open_context.open_pinned() + candidates = _candidate_fragment_ids(pinned, predicate) + if candidates is None: + fragment_ids = [fragment.fragment_id for fragment in pinned.get_fragments()] + logger.info("No scalar index covers %r; running delete over all %d fragments", predicate, len(fragment_ids)) + else: + fragment_ids = sorted(candidates) + logger.info("Scalar index pruned delete for %r down to %d fragments", predicate, len(fragment_ids)) + + updated_fragments, removed_fragment_ids = _delete_matching_rows(open_context, predicate, fragment_ids) + + if not updated_fragments and not removed_fragment_ids and not new_fragments: + return None + + operation = lance.LanceOperation.Update( + removed_fragment_ids=removed_fragment_ids, + updated_fragments=updated_fragments, + new_fragments=list(new_fragments), + # Deletions do not change any field's values, so no index needs to be + # dropped from the fragments that survive. + fields_modified=[], + ) + return lance.LanceDataset.commit( + open_context.uri, + operation, + read_version=open_context.version, + storage_options=open_context.storage_options, + **open_context.commit_kwargs, + ) diff --git a/tests/io/lancedb/test_insert_overwrite.py b/tests/io/lancedb/test_insert_overwrite.py new file mode 100644 index 0000000..b698bed --- /dev/null +++ b/tests/io/lancedb/test_insert_overwrite.py @@ -0,0 +1,318 @@ +"""Conditional overwrite: ``mode="insert_overwrite"`` replaces a predicate's rows in one commit.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import lance +import pyarrow as pa +import pytest + +import daft +import daft_lance +from daft.recordbatch import MicroPartition +from daft_lance.lance_data_sink import LanceDataSink +from daft_lance.lance_insert_overwrite import _SCALAR_INDEX_PLAN_MARKER, _candidate_fragment_ids + + +def _seed(uri: str) -> None: + """Three fragments; the first and second each mix two ``dt`` values.""" + daft_lance.write_lance( + daft.from_pydict({"dt": ["d1", "d2"], "id": [1, 2]}), uri, mode="create", max_rows_per_file=2 + ).collect() + daft_lance.write_lance( + daft.from_pydict({"dt": ["d2", "d3"], "id": [3, 4]}), uri, mode="append", max_rows_per_file=2 + ).collect() + daft_lance.write_lance( + daft.from_pydict({"dt": ["d2", "d2"], "id": [5, 6]}), uri, mode="append", max_rows_per_file=2 + ).collect() + + +def _rows(uri: str) -> list[tuple[str, int]]: + table = lance.dataset(uri).to_table().to_pydict() + return sorted(zip(table["dt"], table["id"])) + + +def _overwrite( + uri: str, dts: list[str | None], ids: list[int], overwrite_where: str, **kwargs: Any +) -> dict[str, list[Any]]: + return daft_lance.write_lance( + daft.from_pydict({"dt": dts, "id": ids}), + uri, + mode="insert_overwrite", + overwrite_where=overwrite_where, + **kwargs, + ).to_pydict() + + +def test_replaces_only_matching_rows_in_one_version(tmp_path: Path) -> None: + uri = str(tmp_path / "tbl") + _seed(uri) + before = lance.dataset(uri).version + + stats = _overwrite(uri, ["d2", "d2"], [100, 101], "dt = 'd2'") + + assert _rows(uri) == [("d1", 1), ("d2", 100), ("d2", 101), ("d3", 4)] + # One commit, not a delete followed by an append: readers never see the gap. + assert lance.dataset(uri).version == before + 1 + assert stats["version"] == [before + 1] + + +def test_rerunning_the_same_overwrite_is_idempotent(tmp_path: Path) -> None: + uri = str(tmp_path / "tbl") + _seed(uri) + expected = [("d1", 1), ("d2", 100), ("d2", 101), ("d3", 4)] + + _overwrite(uri, ["d2", "d2"], [100, 101], "dt = 'd2'") + assert _rows(uri) == expected + # The second run has to delete the rows the first one appended, which sit in + # a fragment that did not exist when the first run planned its delete. + _overwrite(uri, ["d2", "d2"], [100, 101], "dt = 'd2'") + + assert _rows(uri) == expected + + +def test_predicate_matching_nothing_only_appends(tmp_path: Path) -> None: + uri = str(tmp_path / "tbl") + _seed(uri) + before = lance.dataset(uri).version + + _overwrite(uri, ["d9"], [42], "dt = 'd9'") + + assert _rows(uri) == [("d1", 1), ("d2", 2), ("d2", 3), ("d2", 5), ("d2", 6), ("d3", 4), ("d9", 42)] + assert lance.dataset(uri).version == before + 1 + + +def test_empty_input_deletes_the_matched_rows(tmp_path: Path) -> None: + uri = str(tmp_path / "tbl") + _seed(uri) + before = lance.dataset(uri).version + empty = daft.from_pydict({"dt": ["d2"], "id": [1]}).limit(0) + + daft_lance.write_lance(empty, uri, mode="insert_overwrite", overwrite_where="dt = 'd2'").collect() + + assert _rows(uri) == [("d1", 1), ("d3", 4)] + assert lance.dataset(uri).version == before + 1 + + +def test_fully_matched_fragment_is_removed_not_just_emptied(tmp_path: Path) -> None: + uri = str(tmp_path / "tbl") + _seed(uri) + # The third seeded fragment is all "d2", so the overwrite drops it entirely + # while the mixed fragments only gain deletion files. + _overwrite(uri, ["d2"], [100], "dt = 'd2'") + + fragments = lance.dataset(uri).get_fragments() + assert sum(fragment.count_rows() for fragment in fragments) == 3 + assert all(fragment.count_rows() > 0 for fragment in fragments) + + +def test_rows_outside_overwrite_where_are_appended(tmp_path: Path) -> None: + uri = str(tmp_path / "tbl") + _seed(uri) + + _overwrite(uri, ["d2", "d9"], [100, 101], "dt = 'd2'") + + assert _rows(uri) == [("d1", 1), ("d2", 100), ("d3", 4), ("d9", 101)] + + +@pytest.mark.parametrize( + ("kwargs", "match"), + [ + ({"mode": "insert_overwrite"}, "requires a non-empty SQL predicate"), + ({"mode": "insert_overwrite", "overwrite_where": " "}, "requires a non-empty SQL predicate"), + ({"mode": "append", "overwrite_where": "dt = 'd2'"}, 'only supported with mode="insert_overwrite"'), + ( + {"mode": "insert_overwrite", "overwrite_where": "dt = 'd2'", "use_mem_wal": True}, + "not supported with use_mem_wal", + ), + ], +) +def test_argument_validation(tmp_path: Path, kwargs: dict[str, Any], match: str) -> None: + with pytest.raises(ValueError, match=match): + daft_lance.write_lance(daft.from_pydict({"dt": ["d2"], "id": [1]}), str(tmp_path / "tbl"), **kwargs) + + +def test_requires_an_existing_table(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="Cannot insert_overwrite to non-existent Lance dataset"): + _overwrite(str(tmp_path / "missing"), ["d2"], [1], "dt = 'd2'") + + +def test_schema_must_match_like_append(tmp_path: Path) -> None: + uri = str(tmp_path / "tbl") + _seed(uri) + + with pytest.raises(ValueError, match="Schema of data does not match table schema"): + daft_lance.write_lance( + daft.from_pydict({"dt": ["d2"]}), uri, mode="insert_overwrite", overwrite_where="dt = 'd2'" + ).collect() + + +def test_storage_version_conflict_is_detected_like_append(tmp_path: Path) -> None: + """Regression guard for the mode normalization. + + ``resolve_storage_version`` only checks the "append" mode; before + ``insert_overwrite`` was normalized to it, a conflicting version was accepted + silently. + """ + uri = str(tmp_path / "tbl") + daft_lance.write_lance( + daft.from_pydict({"dt": ["d2"], "id": [1]}), uri, mode="create", data_storage_version="2.1" + ).collect() + + with pytest.raises(ValueError, match="does not match existing dataset version"): + _overwrite(uri, ["d2"], [2], "dt = 'd2'", data_storage_version="2.0") + + +@pytest.mark.parametrize("predicate", ["nosuchcol = 1", "dt ==== 'x'"]) +def test_predicate_must_be_a_valid_lance_filter(tmp_path: Path, predicate: str) -> None: + """Bad predicates fail on the driver, before any data is written.""" + uri = str(tmp_path / "tbl") + _seed(uri) + + with pytest.raises(ValueError, match="is not a valid Lance filter"): + _overwrite(uri, ["d2"], [100], predicate) + + +def test_pruning_uses_a_scalar_index_when_one_covers_the_predicate(tmp_path: Path) -> None: + uri = str(tmp_path / "tbl") + _seed(uri) + dataset = lance.dataset(uri) + + # No index: every fragment has to be visited, and the planner says so. + assert _candidate_fragment_ids(dataset, "id = 4") is None + + dataset.create_scalar_index("id", "BTREE") + dataset = lance.dataset(uri) + + # Pinned plan-node name: pruning silently stops working if Lance renames it. + plan = dataset.scanner(columns=[], filter="id = 4", with_row_address=True).explain_plan(True) + assert _SCALAR_INDEX_PLAN_MARKER in plan + + # id 4 lives only in the second seeded fragment. + assert _candidate_fragment_ids(dataset, "id = 4") == {1} + assert _candidate_fragment_ids(dataset, "id > 2") == {1, 2} + assert _candidate_fragment_ids(dataset, "id = 999") == set() + + +def test_overwrite_through_a_scalar_index_on_the_predicate_column(tmp_path: Path) -> None: + """The pruning path, end to end -- a miss here silently leaves rows behind.""" + uri = str(tmp_path / "tbl") + daft_lance.write_lance( + daft.from_pydict({"day": [1, 1, 2, 2, 3], "id": [1, 2, 3, 4, 5]}), uri, mode="create", max_rows_per_file=2 + ).collect() + lance.dataset(uri).create_scalar_index("day", "BTREE") + + def rows() -> list[tuple[int, int]]: + table = lance.dataset(uri).to_table().to_pydict() + return sorted(zip(table["day"], table["id"])) + + def overwrite(new_id: int) -> None: + daft_lance.write_lance( + daft.from_pydict({"day": [2], "id": [new_id]}), + uri, + mode="insert_overwrite", + overwrite_where="day = 2", + ).collect() + + assert _candidate_fragment_ids(lance.dataset(uri), "day = 2") is not None, "expected the pruning path" + overwrite(100) + assert rows() == [(1, 1), (1, 2), (2, 100), (3, 5)] + + # The index does not cover the fragment the first overwrite appended, so this + # run only replaces its row if pruning still finds that fragment. + assert _candidate_fragment_ids(lance.dataset(uri), "day = 2") is not None, "expected the pruning path" + overwrite(200) + assert rows() == [(1, 1), (1, 2), (2, 200), (3, 5)] + + +def test_indexed_table_stays_queryable_after_overwrite(tmp_path: Path) -> None: + uri = str(tmp_path / "tbl") + n = 300 + vector_type = pa.list_(pa.float32(), 2) + seed = pa.table( + { + "id": pa.array(range(n), pa.int64()), + "dt": pa.array(["d1" if i % 2 else "d2" for i in range(n)]), + "vector": pa.array([[float(i % 3), 0.0] for i in range(n)], type=vector_type), + } + ) + lance.write_dataset(seed, uri, max_rows_per_file=100) + dataset = lance.dataset(uri) + dataset.create_scalar_index("id", "BTREE") + try: + dataset.create_index("vector", "IVF_PQ", num_partitions=2, num_sub_vectors=1) + except Exception: + pytest.skip("Could not create vector index (lance version or dataset size issue)") + + new_rows = pa.table( + { + "id": pa.array([1000, 1001], pa.int64()), + "dt": pa.array(["d2", "d2"]), + "vector": pa.array([[7.0, 7.0], [8.0, 8.0]], type=vector_type), + } + ) + daft_lance.write_lance( + daft.from_arrow(new_rows), uri, mode="insert_overwrite", overwrite_where="dt = 'd2'" + ).collect() + + dataset = lance.dataset(uri) + # Deleted rows are invisible through the scalar index that still covers them. + assert dataset.to_table(filter="id = 0").num_rows == 0 + assert dataset.to_table(filter="id = 1").num_rows == 1 + assert sorted(dataset.to_table(filter="dt = 'd2'").to_pydict()["id"]) == [1000, 1001] + + # New fragments are not in the index; the search must still find them. + nearest = {"column": "vector", "q": pa.array([8.0, 8.0], type=pa.float32()), "k": 1, "use_index": True} + assert daft.read_lance(uri, default_scan_options={"nearest": nearest}).select("id").to_pydict()["id"] == [1001] + + +def test_namespace_addressed_table(tmp_path: Path) -> None: + ns: dict[str, Any] = {"namespace_impl": "dir", "namespace_properties": {"root": str(tmp_path)}} + table_id = ["events"] + + daft_lance.write_lance( + daft.from_pydict({"dt": ["d1", "d2"], "id": [1, 2]}), table_id=table_id, mode="create", **ns + ).collect() + daft_lance.write_lance( + daft.from_pydict({"dt": ["d2"], "id": [100]}), + table_id=table_id, + mode="insert_overwrite", + overwrite_where="dt = 'd2'", + **ns, + ).collect() + + result = daft_lance.read_lance(table_id=table_id, **ns).to_pydict() + assert sorted(zip(result["dt"], result["id"])) == [("d1", 1), ("d2", 100)] + + +def test_concurrent_append_survives_the_overwrite(tmp_path: Path) -> None: + """Documents a known gap, so a future Lance change surfaces here. + + Lance does not treat a concurrent append as conflicting with the Update this + mode commits, so rows another writer adds while the overwrite runs stay in + the table even when they match the predicate -- and nothing raises. The + docstring on ``write_lance`` warns about it; this test pins the behavior. + """ + uri = str(tmp_path / "tbl") + _seed(uri) + + sink = LanceDataSink( + uri=uri, + schema=daft.from_pydict({"dt": ["d2"], "id": [1]}).schema(), + mode="insert_overwrite", + overwrite_where="dt = 'd2'", + ) + sink.start() # pins the read version + + concurrent = pa.table({"dt": pa.array(["d2"], pa.large_string()), "id": pa.array([900], pa.int64())}) + lance.write_dataset(concurrent, uri, mode="append") + + results = list(sink.write(iter([MicroPartition.from_pydict({"dt": ["d2"], "id": [100]})]))) + sink.finalize(results) + + # The overwrite replaced every "d2" row it knew about, and the concurrent one + # it could not see survived: asserting the whole table keeps this test honest + # if the overwrite ever silently turns into a no-op. + assert _rows(uri) == [("d1", 1), ("d2", 100), ("d2", 900), ("d3", 4)]