Skip to content
Open
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
61 changes: 55 additions & 6 deletions daft_lance/lance_scan.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
from __future__ import annotations

import logging
from collections.abc import Iterator
from collections.abc import Callable, Iterator
from typing import Any

import lance
Expand All @@ -23,6 +23,47 @@
logger = logging.getLogger(__name__)


class _LanceBatchIterator(Iterator[PyRecordBatch]):
"""Iterator over ``PyRecordBatch`` that also exposes cumulative I/O statistics.

Daft's executor duck-types the object returned by a ``python_factory_func_scan_task``
factory: if it has a callable ``stats`` attribute, Daft calls it after every batch and
once more when the iterator is exhausted, and folds the *delta* since the previous call
into the scan task's IOStats. ``stats()`` therefore returns cumulative counters for the
lifetime of this iterator, and the counters only ever increase.

Every ``ds.scanner(...)`` created by the wrapped generator must be given
``scan_stats_callback=self.record_scan_stats`` so that Lance's per-scanner
``ScanStatistics`` are added into the running totals.
"""

__slots__ = ("_batches", "_bytes_read", "_requests")

def __init__(
self, make_batches: Callable[[Callable[[lance.ScanStatistics], None]], Iterator[PyRecordBatch]]
) -> None:
self._bytes_read = 0
self._requests = 0
# ``make_batches`` receives this wrapper's callback and must pass it as
# ``scan_stats_callback`` to every scanner it creates.
self._batches = make_batches(self.record_scan_stats)

def record_scan_stats(self, scan_stats: lance.ScanStatistics) -> None:
"""Lance ``scan_stats_callback``; fires once per scanner when it finishes."""
self._bytes_read += int(getattr(scan_stats, "bytes_read", 0) or 0)
self._requests += int(getattr(scan_stats, "requests", 0) or 0)

def stats(self) -> dict[str, int]:
"""Cumulative I/O counters for this iterator, in the keys Daft recognizes."""
return {"bytes.read": self._bytes_read, "io.requests": self._requests}

def __iter__(self) -> _LanceBatchIterator:
return self

def __next__(self) -> PyRecordBatch:
return next(self._batches)


# TODO support fts and fast_search
def _lancedb_table_factory_function(
ds_uri: str,
Expand All @@ -33,7 +74,7 @@ def _lancedb_table_factory_function(
limit: int | None = None,
include_fragment_id: bool | None = False,
nearest: dict[str, Any] | None = None,
) -> Iterator[PyRecordBatch]:
) -> _LanceBatchIterator:
if fragment_ids is not None and nearest is not None:
raise ValueError(
"fragment_ids and nearest options are mutually exclusive. "
Expand All @@ -43,7 +84,9 @@ def _lancedb_table_factory_function(

ds = open_dataset_from_open_kwargs(ds_uri, open_kwargs)

def _iter_batches() -> Iterator[PyRecordBatch]:
def _iter_batches(
fragments: list[lance.LanceFragment], record_scan_stats: Callable[[lance.ScanStatistics], None]
) -> Iterator[PyRecordBatch]:
# Iterate fragments individually; append a fragment_id column only when requested
# Handle limit correctly by tracking how many rows we've yielded so far
rows_yielded = 0
Expand All @@ -66,6 +109,7 @@ def _iter_batches() -> Iterator[PyRecordBatch]:
filter=filter,
limit=fragment_limit,
blob_handling="blobs_descriptions",
scan_stats_callback=record_scan_stats,
)

for rb in scanner.to_batches():
Expand All @@ -88,27 +132,32 @@ def _iter_batches() -> Iterator[PyRecordBatch]:
yield RecordBatch.from_arrow_record_batches([rb], rb.schema)._recordbatch
rows_yielded += len(rb)

# If fragment_ids is None, let Lance choose fragments via index; omit the fragments parameter.
if fragment_ids is None:
def _index_driven_batches(record_scan_stats: Callable[[lance.ScanStatistics], None]) -> Iterator[PyRecordBatch]:
# Let Lance choose fragments via index; omit the fragments parameter.
# The scanner is built eagerly (as before) so construction errors surface at factory-call time.
scanner = ds.scanner(
columns=required_columns,
filter=filter,
limit=limit,
nearest=nearest,
blob_handling="blobs_descriptions",
scan_stats_callback=record_scan_stats,
)

def _batches() -> Iterator[PyRecordBatch]:
for rb in scanner.to_batches():
yield RecordBatch.from_arrow_record_batches([rb], rb.schema)._recordbatch

return _batches()

if fragment_ids is None:
return _LanceBatchIterator(_index_driven_batches)
else:
fragments_raw = [ds.get_fragment(id) for id in (fragment_ids or [])]
fragments = [f for f in fragments_raw if f is not None]
if not fragments:
raise RuntimeError(f"Unable to find lance fragments {fragment_ids}")
return _iter_batches()
return _LanceBatchIterator(lambda record_scan_stats: _iter_batches(fragments, record_scan_stats))


def _lancedb_count_result_function(
Expand Down
247 changes: 247 additions & 0 deletions tests/io/lancedb/test_lancedb_scan_stats.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,247 @@
from __future__ import annotations

from collections.abc import Iterator
from pathlib import Path
from typing import Any

import pyarrow as pa
import pyarrow.compute as pc
import pytest

import daft
from daft.subscribers import Subscriber
from daft.subscribers.events import Stats
from daft_lance.lance_scan import _LanceBatchIterator, _lancedb_table_factory_function

lance = pytest.importorskip("lance")

NUM_FRAGMENTS = 3
ROWS_PER_FRAGMENT = 500


@pytest.fixture(scope="function")
def lance_dataset_path(tmp_path: Path) -> str:
for frag_idx in range(NUM_FRAGMENTS):
base = frag_idx * ROWS_PER_FRAGMENT
tbl = pa.Table.from_pydict(
{
"big_int": list(range(base, base + ROWS_PER_FRAGMENT)),
"payload": [f"row-{i:06d}" * 4 for i in range(base, base + ROWS_PER_FRAGMENT)],
}
)
lance.write_dataset(tbl, tmp_path, mode="append" if frag_idx > 0 else None)
return str(tmp_path)


def _drain_and_snapshot(it: _LanceBatchIterator) -> tuple[int, list[dict[str, int]]]:
"""Drain the iterator the way Daft does: poll ``stats()`` after every batch and once at the end."""
snapshots: list[dict[str, int]] = []
rows = 0
for rb in it:
rows += len(rb)
snapshots.append(it.stats())
snapshots.append(it.stats())
return rows, snapshots


def _assert_monotonic(snapshots: list[dict[str, int]]) -> None:
for prev, cur in zip(snapshots, snapshots[1:]):
assert cur["bytes.read"] >= prev["bytes.read"]
assert cur["io.requests"] >= prev["io.requests"]


# ---------------------------------------------------------------------------
# Unit tests: call the factory directly (no dependency on the Daft executor).
# ---------------------------------------------------------------------------


def test_factory_returns_iterator_with_stats(lance_dataset_path: str) -> None:
it = _lancedb_table_factory_function(ds_uri=lance_dataset_path, fragment_ids=[0])
assert isinstance(it, Iterator)
assert iter(it) is it
assert callable(it.stats)
# Nothing has been read yet.
assert it.stats() == {"bytes.read": 0, "io.requests": 0}


def test_factory_stats_per_fragment_path(lance_dataset_path: str) -> None:
ds = lance.dataset(lance_dataset_path)
frag_ids = [f.fragment_id for f in ds.get_fragments()]
it = _lancedb_table_factory_function(ds_uri=lance_dataset_path, fragment_ids=frag_ids)
rows, snapshots = _drain_and_snapshot(it)

assert rows == NUM_FRAGMENTS * ROWS_PER_FRAGMENT
final = snapshots[-1]
assert set(final) == {"bytes.read", "io.requests"}
assert final["bytes.read"] > 0
assert final["io.requests"] > 0
_assert_monotonic(snapshots)
# Cumulative: repeated polling after exhaustion does not grow the counters.
assert it.stats() == final
assert it.stats() == final


def test_factory_stats_index_driven_path(lance_dataset_path: str) -> None:
it = _lancedb_table_factory_function(ds_uri=lance_dataset_path, fragment_ids=None)
rows, snapshots = _drain_and_snapshot(it)

assert rows == NUM_FRAGMENTS * ROWS_PER_FRAGMENT
final = snapshots[-1]
assert final["bytes.read"] > 0
assert final["io.requests"] > 0
_assert_monotonic(snapshots)
assert it.stats() == final


def test_factory_stats_index_driven_with_filter_and_limit(lance_dataset_path: str) -> None:
it = _lancedb_table_factory_function(
ds_uri=lance_dataset_path,
fragment_ids=None,
required_columns=["big_int"],
filter=pc.greater_equal(pc.field("big_int"), pc.scalar(ROWS_PER_FRAGMENT)),
limit=7,
)
rows, snapshots = _drain_and_snapshot(it)
assert rows == 7
assert snapshots[-1]["bytes.read"] > 0
assert it.stats() == snapshots[-1]


def test_factory_stats_per_fragment_with_limit_stops_early(lance_dataset_path: str) -> None:
"""The limit stops iteration after the first fragment, so only that scanner contributes stats."""
ds = lance.dataset(lance_dataset_path)
frag_ids = [f.fragment_id for f in ds.get_fragments()]

limited = _lancedb_table_factory_function(ds_uri=lance_dataset_path, fragment_ids=frag_ids, limit=10)
rows, snapshots = _drain_and_snapshot(limited)
assert rows == 10
limited_final = snapshots[-1]
assert limited_final["bytes.read"] > 0

full = _lancedb_table_factory_function(ds_uri=lance_dataset_path, fragment_ids=frag_ids)
_, full_snapshots = _drain_and_snapshot(full)
assert full_snapshots[-1]["bytes.read"] > limited_final["bytes.read"]
assert full_snapshots[-1]["io.requests"] > limited_final["io.requests"]


def test_factory_stats_include_fragment_id(lance_dataset_path: str) -> None:
it = _lancedb_table_factory_function(ds_uri=lance_dataset_path, fragment_ids=[0, 1], include_fragment_id=True)
batches = list(it)
assert all("fragment_id" in rb.schema().names() for rb in batches)
assert it.stats()["bytes.read"] > 0


def test_batch_iterator_accumulates_callback_values() -> None:
"""Pure unit test of the wrapper: callback values are summed, and stats() is cumulative."""

class _FakeScanStatistics:
def __init__(self, bytes_read: int, requests: int) -> None:
self.bytes_read = bytes_read
self.requests = requests
self.iops = requests

def _make_batches(record: Any) -> Iterator[Any]:
record(_FakeScanStatistics(100, 2))
yield "batch-1"
record(_FakeScanStatistics(50, 1))
yield "batch-2"
record(_FakeScanStatistics(0, 0))

it: Any = _LanceBatchIterator(_make_batches)
assert it.stats() == {"bytes.read": 0, "io.requests": 0}
assert next(it) == "batch-1"
assert it.stats() == {"bytes.read": 100, "io.requests": 2}
assert next(it) == "batch-2"
assert it.stats() == {"bytes.read": 150, "io.requests": 3}
with pytest.raises(StopIteration):
next(it)
assert it.stats() == {"bytes.read": 150, "io.requests": 3}
assert it.stats() == {"bytes.read": 150, "io.requests": 3}


# ---------------------------------------------------------------------------
# End-to-end: Daft's executor polls ``stats()`` and surfaces it as ``bytes.read``
# on the scan node. Only supported by Daft builds that expose DataSourceTask.stats().
# ---------------------------------------------------------------------------

_DAFT_SUPPORTS_FACTORY_STATS = hasattr(getattr(daft.io, "DataSourceTask", None), "stats")


class _StatsRecorder(Subscriber):
def __init__(self) -> None:
self.stats_events: list[Stats] = []
self.scan_node_ids: set[int] = set()

def on_operator_start(self, event: Any) -> None:
if "lance" in event.name.lower() or "scan" in event.name.lower():
self.scan_node_ids.add(event.node_id)

def on_stats(self, event: Stats) -> None:
self.stats_events.append(event)

def scan_bytes_read(self) -> list[int]:
"""``bytes.read`` reported for scan nodes, in the order the Stats events arrived."""
values: list[int] = []
for event in self.stats_events:
for node_id, node_stats in event.stats.items():
if node_id not in self.scan_node_ids:
continue
if "bytes.read" in node_stats:
values.append(int(node_stats["bytes.read"][1]))
return values


def _run_with_recorder(df: daft.DataFrame) -> _StatsRecorder:
recorder = _StatsRecorder()
alias = f"lance-scan-stats-{id(recorder)}"
daft.attach_subscriber(alias, recorder)
try:
df.collect()
finally:
daft.detach_subscriber(alias)
return recorder


@pytest.mark.skipif(not _DAFT_SUPPORTS_FACTORY_STATS, reason="Daft build does not support factory iterator stats()")
def test_scan_node_reports_bytes_read_per_fragment_path(lance_dataset_path: str) -> None:
df = daft.read_lance(lance_dataset_path)
recorder = _run_with_recorder(df)

values = recorder.scan_bytes_read()
assert values, "no bytes.read stat was emitted for the scan node"
assert max(values) > 0
# Stats are cumulative on the iterator, so Daft's delta folding must not re-add them:
# every Stats event reports a value no larger than the final one, and the final value
# matches what Lance reports for a full scan of the dataset.
final = values[-1]
assert all(v <= final for v in values)

expected = 0

def _record(st: Any) -> None:
nonlocal expected
expected += st.bytes_read

ds = lance.dataset(lance_dataset_path)
for fragment in ds.get_fragments():
list(
ds.scanner(
fragments=[fragment], scan_stats_callback=_record, blob_handling="blobs_descriptions"
).to_batches()
)
assert final == expected


@pytest.mark.skipif(not _DAFT_SUPPORTS_FACTORY_STATS, reason="Daft build does not support factory iterator stats()")
def test_scan_node_reports_bytes_read_index_driven_path(lance_dataset_path: str) -> None:
ds = lance.dataset(lance_dataset_path)
ds.create_scalar_index("big_int", index_type="BTREE")

df = daft.read_lance(lance_dataset_path).where(daft.col("big_int") == 42)
recorder = _run_with_recorder(df)

values = recorder.scan_bytes_read()
assert values, "no bytes.read stat was emitted for the scan node"
final = values[-1]
assert final > 0
assert all(v <= final for v in values)