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
57 changes: 37 additions & 20 deletions vortex-python/src/dataset.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,15 +6,16 @@ use std::sync::Arc;
use arrow_array::RecordBatchReader;
use arrow_schema::SchemaRef;
use itertools::Itertools;
use pyo3::exceptions::PyIndexError;
use pyo3::exceptions::PyTypeError;
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use pyo3::types::PyString;
use vortex::array::ArrayRef;
use vortex::array::ExecutionCtx;
use vortex::array::VortexSessionExecute;
use vortex::array::arrays::PrimitiveArray;
use vortex::array::iter::ArrayIteratorExt;
use vortex::dtype::DType;
use vortex::dtype::FieldName;
use vortex::dtype::FieldNames;
use vortex::error::VortexResult;
Expand Down Expand Up @@ -83,24 +84,40 @@ pub fn read_array_from_reader(
scan.into_array_iter(&runtime)?.read_all()
}

fn projection_from_python(columns: Option<Vec<Bound<PyAny>>>) -> PyResult<Expression> {
fn field_from_pyany(field: &Bound<PyAny>) -> PyResult<FieldName> {
if field.clone().is_instance_of::<PyString>() {
Ok(FieldName::from(field.cast::<PyString>()?.to_str()?))
} else {
Err(PyTypeError::new_err(format!(
"projection: expected list of strings or None, but found: {field}.",
)))
}
}
/// A projected column, selected either by name or by positional index.
#[derive(FromPyObject)]
pub enum ProjectionColumn {
Name(String),
Index(usize),
}

fn projection_from_python(
columns: Option<Vec<ProjectionColumn>>,
dtype: &DType,
) -> PyResult<Expression> {
Ok(match columns {
None => root(),
Some(columns) => {
let fields: Vec<_> = columns
.iter()
.map(field_from_pyany)
.collect::<PyResult<_>>()?;
let fields = columns
.into_iter()
.map(|column| match column {
ProjectionColumn::Name(name) => Ok(FieldName::from(name.as_str())),
ProjectionColumn::Index(index) => {
// Positional projection: map the index onto the top-level field name.
let DType::Struct(struct_dtype, _) = dtype else {
return Err(PyTypeError::new_err(
"projection: integer indices are only valid for a struct-typed file",
));
};
struct_dtype.field_name(index).cloned().ok_or_else(|| {
PyIndexError::new_err(format!(
"projection: column index {index} is out of range for {} columns",
struct_dtype.nfields()
))
})
}
})
.collect::<PyResult<Vec<_>>>()?;
select(FieldNames::from(fields), root())
}
})
Expand Down Expand Up @@ -142,13 +159,13 @@ impl PyVortexDataset {
pub(crate) fn to_array_inner<'py>(
&self,
py: Python<'py>,
columns: Option<Vec<Bound<'py, PyAny>>>,
columns: Option<Vec<ProjectionColumn>>,
row_filter: Option<&Bound<'py, PyExpr>>,
indices: Option<PyArrayRef>,
row_range: Option<(u64, u64)>,
) -> PyVortexResult<PyArrayRef> {
let vxf = self.vxf.clone();
let projection = projection_from_python(columns)?;
let projection = projection_from_python(columns, vxf.dtype())?;
let filter = filter_from_python(row_filter);
let indices = indices.map(|i| i.into_inner());

Expand All @@ -170,7 +187,7 @@ impl PyVortexDataset {
#[pyo3(signature = (*, columns = None, row_filter = None, indices = None, row_range = None))]
pub fn to_array<'py>(
self_: PyRef<'py, Self>,
columns: Option<Vec<Bound<'py, PyAny>>>,
columns: Option<Vec<ProjectionColumn>>,
row_filter: Option<&Bound<'py, PyExpr>>,
indices: Option<PyArrayRef>,
row_range: Option<(u64, u64)>,
Expand All @@ -181,13 +198,13 @@ impl PyVortexDataset {
#[pyo3(signature = (*, columns = None, row_filter = None, split_by = None, row_range = None))]
pub fn to_record_batch_reader(
self_: PyRef<Self>,
columns: Option<Vec<Bound<'_, PyAny>>>,
columns: Option<Vec<ProjectionColumn>>,
row_filter: Option<&Bound<'_, PyExpr>>,
split_by: Option<usize>,
row_range: Option<(u64, u64)>,
) -> PyVortexResult<Py<PyAny>> {
let vxf = self_.vxf.clone();
let projection = projection_from_python(columns)?;
let projection = projection_from_python(columns, vxf.dtype())?;
let filter = filter_from_python(row_filter);

let reader = self_.py().detach(move || {
Expand Down
3 changes: 2 additions & 1 deletion vortex-python/src/io.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ use crate::arrow::FromPyArrow;
use crate::classes::record_batch_reader_class;
use crate::classes::table_class;
use crate::current_runtime;
use crate::dataset::ProjectionColumn;
use crate::dataset::PyVortexDataset;
use crate::error::PyVortexResult;
use crate::expr::PyExpr;
Expand Down Expand Up @@ -125,7 +126,7 @@ pub fn read_url<'py>(
py: Python<'py>,
url: &str,
store: Option<Bound<'py, PyAny>>,
projection: Option<Vec<Bound<'py, PyAny>>>,
projection: Option<Vec<ProjectionColumn>>,
row_filter: Option<&Bound<'py, PyExpr>>,
indices: Option<PyArrayRef>,
row_range: Option<(u64, u64)>,
Expand Down
23 changes: 23 additions & 0 deletions vortex-python/test/test_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,3 +40,26 @@ def test_store_roundtrip(tmp_path: Path) -> None:
people = vx.io.read_url("people.vortex", store=local)

assert people.to_pylist() == records.to_pylist()


def test_read_url_integer_projection(tmp_path: Path) -> None:
local = LocalStore(prefix=tmp_path)
records = vx.array([dict(name="Alice", salary=10), dict(name="Bob", salary=20)])
vx.io.write(records, "people.vortex", store=local)

# Columns are name (0) and salary (1); select salary by position.
by_index = vx.io.read_url("people.vortex", store=local, projection=[1])
assert by_index.to_pylist() == [{"salary": 10}, {"salary": 20}]

# Integer and name projection agree.
by_name = vx.io.read_url("people.vortex", store=local, projection=["salary"])
assert by_index.to_pylist() == by_name.to_pylist()


def test_read_url_integer_projection_out_of_range(tmp_path: Path) -> None:
local = LocalStore(prefix=tmp_path)
records = vx.array([dict(name="Alice", salary=10)])
vx.io.write(records, "people.vortex", store=local)

with pytest.raises(IndexError):
vx.io.read_url("people.vortex", store=local, projection=[99])
Loading