Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 25 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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/)
Expand Down
19 changes: 17 additions & 2 deletions daft_lance/_lance.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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.
Expand All @@ -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)

Expand All @@ -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,
Expand Down
153 changes: 124 additions & 29 deletions daft_lance/lance_data_sink.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand All @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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,
Expand All @@ -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]
Expand All @@ -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.
Expand Down Expand Up @@ -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.")
Expand All @@ -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),
):
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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":
Expand All @@ -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)
Expand Down Expand Up @@ -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:
Expand Down
Loading
Loading