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
1 change: 1 addition & 0 deletions vortex-python/python/vortex/_lib/file.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@ class VortexFile:
*,
expr: Expr | None = None,
limit: int | None = None,
indices: Array | None = None,
batch_size: int | None = None,
schema: pa.Schema | None = None,
) -> pa.RecordBatchReader: ...
Expand Down
7 changes: 6 additions & 1 deletion vortex-python/python/vortex/file.py
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,7 @@ def to_arrow(
*,
limit: int | None = None,
expr: Expr | None = None,
indices: Array | None = None,
batch_size: int | None = None,
schema: pa.Schema | None = None,
) -> RecordBatchReader:
Expand All @@ -220,14 +221,18 @@ def to_arrow(
from the file) or an explicit list of desired columns.
expr : :class:`vortex.Expr` | None
The predicate used to filter rows. The filter columns need not appear in the projection.
indices : :class:`vortex.Array` | None
The indices of the rows to read. Must be strictly increasing and non-null.
batch_size : :class:`int` | None
The number of rows to read per chunk.
schema : :class:`pyarrow.Schema` | None
The Arrow schema to return. Use ``pyarrow.string()`` for ``StringArray`` fields.
Use ``pyarrow.binary()`` for ``BinaryArray`` fields.

"""
return self._file.to_arrow(projection, expr=expr, limit=limit, batch_size=batch_size, schema=schema)
return self._file.to_arrow(
projection, expr=expr, limit=limit, indices=indices, batch_size=batch_size, schema=schema
)

def to_dataset(self) -> VortexDataset:
"""Scan the Vortex file using the :class:`pyarrow.dataset.Dataset` API."""
Expand Down
27 changes: 6 additions & 21 deletions vortex-python/src/file.rs
Original file line number Diff line number Diff line change
Expand Up @@ -201,43 +201,28 @@ impl PyVortexFile {
})
}

#[pyo3(signature = (projection = None, *, expr = None, limit = None, batch_size = None, schema = None))]
#[pyo3(signature = (projection = None, *, expr = None, limit = None, indices = None, batch_size = None, schema = None))]
fn to_arrow(
slf: Bound<Self>,
projection: Option<PyIntoProjection>,
expr: Option<PyExpr>,
limit: Option<u64>,
indices: Option<PyArrayRef>,
batch_size: Option<usize>,
schema: Option<&Bound<PyAny>>,
) -> PyVortexResult<Py<PyAny>> {
let vxf = &slf.get().vxf;
let projection = projection.map(|p| p.0);
let expr = expr.map(|e| e.into_inner());
let indices = row_indices(slf.py(), indices)?;
let schema = schema
.map(|schema| Schema::from_pyarrow(&schema.as_borrowed()))
.transpose()?
.map(Arc::new);

// Building the reader is lazy and cheap, so it runs without releasing the GIL. The scan
// runs as pyarrow pulls batches, and pyarrow releases the GIL while it does.
let filter = expr
.map(|e| e.into_inner().bind(vxf.dtype())?.optimize_recursive())
.transpose()?;
let projection = projection
.map(|p| p.0)
.unwrap_or_else(root)
.bind(vxf.dtype())?
.optimize_recursive()?;
let mut builder = vxf
.scan()?
.with_some_filter(filter)
.with_projection(projection);

if let Some(limit) = limit {
builder = builder.with_limit(limit);
}

if let Some(batch_size) = batch_size {
builder = builder.with_split_by(SplitBy::RowCount(batch_size));
}
let builder = scan_builder(vxf, projection, expr, limit, indices, batch_size)?;

let schema = match schema {
Some(schema) => schema,
Expand Down
74 changes: 74 additions & 0 deletions vortex-python/test/test_file_arrow_indices.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright the Vortex contributors

from pathlib import Path

import pyarrow as pa
import pytest

import vortex as vx
import vortex.expr as ve


@pytest.fixture
def indexed_file(tmp_path: Path) -> tuple[vx.VortexFile, pa.Table]:
table = pa.table(
{
"id": list(range(9)),
"payload": ["first", None, "long string with 数据", "", "four", "five", None, "seven", "last"],
}
)
path = tmp_path / "indexed.vortex"
vx.io.write(vx.compress(vx.array(table)), str(path))
return vx.open(str(path)), table


@pytest.mark.parametrize("indices", [[], [1], [0, 2, 5, 8]])
@pytest.mark.parametrize("batch_size", [1, 3])
def test_arrow_indices_preserve_values_and_requested_schema(indexed_file, indices, batch_size):
file, table = indexed_file
indices = pa.array(indices, type=pa.uint64())
schema = pa.schema([("payload", pa.string())])
reader = file.to_arrow(["payload"], indices=vx.array(indices), batch_size=batch_size, schema=schema)
batches = list(reader)
scan_batches = list(file.scan(["payload"], indices=vx.array(indices), batch_size=batch_size))
assert [batch.num_rows for batch in batches] == [len(batch) for batch in scan_batches]
result = pa.Table.from_batches(batches, schema=reader.schema)
assert result.equals(table.select(["payload"]).take(indices))


def test_arrow_indices_apply_filter(indexed_file):
file, table = indexed_file
reader = file.to_arrow(
["payload"],
indices=vx.array([0, 2, 5, 8]),
expr=ve.column("id") >= 4,
schema=pa.schema([("payload", pa.string())]),
)
assert reader.read_all().equals(table.select(["payload"]).take(pa.array([5, 8])))


def test_arrow_indices_apply_limit(indexed_file):
file, table = indexed_file
reader = file.to_arrow(
["payload"],
indices=vx.array([0, 2, 5, 8]),
limit=1,
schema=pa.schema([("payload", pa.string())]),
)
assert reader.read_all().equals(table.select(["payload"]).slice(0, 1))


def test_arrow_indices_use_scan_filter_limit_validation(indexed_file):
file, _ = indexed_file
reader = file.to_arrow(indices=vx.array([0, 2, 5, 8]), expr=ve.column("id") >= 4, limit=1)
with pytest.raises(pa.ArrowInvalid, match="doesn't support scans with both a filter and a limit"):
reader.read_all()


@pytest.mark.parametrize("indices", [[2, 1], [1, 1], [None, 1]])
def test_arrow_indices_use_scan_validation(indexed_file, indices):
file, _ = indexed_file
indices = vx.array(pa.array(indices, type=pa.uint64()))
with pytest.raises(RuntimeError):
file.to_arrow(indices=indices)
Loading