use vortex_buffer::BitBuffer;
use vortex_error::VortexResult;
use vortex_error::vortex_bail;
use vortex_error::vortex_ensure_eq;
use vortex_mask::MaskValuesRef;
use crate::ArrayRef;
use crate::ExecutionCtx;
use crate::scalar_fn::ExecutionArgs;
use crate::scalar_fn::unstable::row::ElementTuple;
use crate::scalar_fn::unstable::row::OutputSink;
use crate::scalar_fn::unstable::row::SinkResult;
use crate::scalar_fn::unstable::row::ViewLen;
pub(crate) fn execute_sink<Args, Prepared, Sink, ApplyResult>(
args: &dyn ExecutionArgs,
params: &Sink::Params,
ctx: &mut ExecutionCtx,
prepare: impl FnOnce(Args::ConstElems<'_>) -> Prepared,
apply: impl Fn(&Prepared, Args::Elems<'_>, Sink::Row<'_>) -> ApplyResult,
) -> VortexResult<ArrayRef>
where
Args: ElementTuple,
Sink: OutputSink,
ApplyResult: SinkResult<WriteToken = Sink::WriteToken>,
{
let columns = Args::decode(args, ctx)?;
let row_count = args.row_count();
let const_values = Args::const_values(&columns);
let prepared = prepare(const_values);
let mut sink = Sink::with_capacity(row_count, params)?;
{
let mut rows = Sink::rows(&mut sink);
let sink_row_count = rows.len();
vortex_ensure_eq!(
sink_row_count,
row_count,
"the output sink must address exactly {row_count} rows, got {sink_row_count}",
);
let views = Args::views_if_no_consts(&columns);
if let Some(views) = views {
if !Args::view_lens_match(&views, row_count) {
decoded_length_error(row_count)?;
}
for index in 0..row_count {
let elements = unsafe { Args::get_from_views_unchecked(&views, index) };
let output = unsafe { Sink::row_unchecked(&mut rows, index) };
apply(&prepared, elements, output).into_result()?;
}
} else {
if !Args::decoded_lens_match(&columns, row_count) {
decoded_length_error(row_count)?;
}
for index in 0..row_count {
let output = unsafe { Sink::row_unchecked(&mut rows, index) };
apply(&prepared, Args::get(&columns, index), output).into_result()?;
}
}
}
unsafe { Sink::finish(sink) }
}
pub(crate) fn execute_sink_valid_rows<Args, Prepared, Sink, ApplyResult>(
args: &dyn ExecutionArgs,
valid: &MaskValuesRef,
params: &Sink::Params,
ctx: &mut ExecutionCtx,
prepare: impl FnOnce(Args::ConstElems<'_>) -> Prepared,
apply: impl Fn(&Prepared, Args::Elems<'_>, Sink::Row<'_>) -> ApplyResult,
) -> VortexResult<Option<ArrayRef>>
where
Args: ElementTuple,
Sink: OutputSink,
ApplyResult: SinkResult<WriteToken = Sink::WriteToken>,
{
let Some(ValidRowsSetup {
columns,
valid_rows,
row_count,
mut sink,
}) = setup_sink_valid_rows::<Args, Sink>(args, valid, params, ctx)?
else {
return Ok(None);
};
let views = Args::views_if_no_consts(&columns);
let const_values = Args::const_values(&columns);
let prepared = prepare(const_values);
{
let mut rows = Sink::rows(&mut sink);
Sink::initialize_skipped_rows(&mut rows);
let initialized_row_count = rows.len();
vortex_ensure_eq!(
initialized_row_count,
row_count,
"the initialized output sink must address exactly {row_count} rows, got {initialized_row_count}",
);
if let Some(views) = views {
if !Args::view_lens_match(&views, row_count) {
decoded_length_error(row_count)?;
}
valid_rows.try_for_each_set_index(|index| {
let output = unsafe { Sink::row_unchecked(&mut rows, index) };
let elements = unsafe { Args::get_from_views_unchecked(&views, index) };
apply(&prepared, elements, output).into_result()
})?;
} else {
if !Args::decoded_lens_match(&columns, row_count) {
decoded_length_error(row_count)?;
}
valid_rows.try_for_each_set_index(|index| {
let output = unsafe { Sink::row_unchecked(&mut rows, index) };
apply(&prepared, Args::get(&columns, index), output).into_result()
})?;
}
}
unsafe { Sink::finish(sink) }.map(Some)
}
pub(crate) fn execute_sink_filtered<Args, Prepared, Sink, ApplyResult>(
args: &dyn ExecutionArgs,
valid: &MaskValuesRef,
params: &Sink::Params,
ctx: &mut ExecutionCtx,
prepare: impl FnOnce(Args::ConstElems<'_>) -> Prepared,
apply: impl Fn(&Prepared, Args::Elems<'_>, Sink::Row<'_>) -> ApplyResult,
) -> VortexResult<ArrayRef>
where
Args: ElementTuple,
Sink: OutputSink,
ApplyResult: SinkResult<WriteToken = Sink::WriteToken>,
{
let columns = Args::decode(args, ctx)?;
let filtered_len = args.row_count();
vortex_ensure_eq!(
valid.true_count(),
filtered_len,
"the filtered batch must contain one row per valid row: {} valid rows, got {filtered_len}",
valid.true_count(),
);
let original_len = valid.len();
let mut sink = Sink::with_capacity(original_len, params)?;
let valid_rows = valid.bit_buffer();
let views = Args::views_if_no_consts(&columns);
let const_values = Args::const_values(&columns);
let prepared = prepare(const_values);
{
let mut rows = Sink::rows(&mut sink);
Sink::initialize_skipped_rows(&mut rows);
let initialized_row_count = rows.len();
vortex_ensure_eq!(
initialized_row_count,
original_len,
"the initialized output sink must address exactly {original_len} rows, got {initialized_row_count}",
);
let mut filtered_index = 0;
if let Some(views) = views {
if !Args::view_lens_match(&views, filtered_len) {
decoded_length_error(filtered_len)?;
}
valid_rows.try_for_each_set_index(|index| {
let output = unsafe { Sink::row_unchecked(&mut rows, index) };
let elements = unsafe { Args::get_from_views_unchecked(&views, filtered_index) };
filtered_index += 1;
apply(&prepared, elements, output).into_result()
})?;
} else {
if !Args::decoded_lens_match(&columns, filtered_len) {
decoded_length_error(filtered_len)?;
}
valid_rows.try_for_each_set_index(|index| {
let output = unsafe { Sink::row_unchecked(&mut rows, index) };
let elements = Args::get(&columns, filtered_index);
filtered_index += 1;
apply(&prepared, elements, output).into_result()
})?;
}
}
unsafe { Sink::finish(sink) }
}
#[cold]
#[inline(never)]
fn decoded_length_error(row_count: usize) -> VortexResult<()> {
vortex_bail!("a decoded row input does not address exactly {row_count} rows")
}
struct ValidRowsSetup<'valid, Args, Sink>
where
Args: ElementTuple,
Sink: OutputSink,
{
columns: Args::Columns,
valid_rows: &'valid BitBuffer,
row_count: usize,
sink: Sink,
}
fn setup_sink_valid_rows<'valid, Args, Sink>(
args: &dyn ExecutionArgs,
valid: &'valid MaskValuesRef,
params: &Sink::Params,
ctx: &mut ExecutionCtx,
) -> VortexResult<Option<ValidRowsSetup<'valid, Args, Sink>>>
where
Args: ElementTuple,
Sink: OutputSink,
{
let Some(columns) = Args::decode_null_tolerant(args, ctx)? else {
return Ok(None);
};
let row_count = args.row_count();
let sink = Sink::with_capacity(row_count, params)?;
let valid_rows = valid.bit_buffer();
vortex_ensure_eq!(
valid_rows.len(),
row_count,
"the validity mask must address exactly {row_count} rows, got {}",
valid_rows.len(),
);
Ok(Some(ValidRowsSetup {
columns,
valid_rows,
row_count,
sink,
}))
}
#[cfg(test)]
mod tests {
use vortex_error::VortexResult;
use vortex_error::vortex_bail;
use vortex_mask::Mask;
use super::execute_sink_valid_rows;
use crate::ArrayRef;
use crate::IntoArray;
use crate::VortexSessionExecute;
use crate::array_session;
use crate::arrays::PrimitiveArray;
use crate::dtype::DType;
use crate::dtype::NativePType;
use crate::scalar_fn::VecExecutionArgs;
use crate::scalar_fn::unstable::row::OutputSink;
struct ShrinkingSink(Vec<i64>);
unsafe impl OutputSink for ShrinkingSink {
type Params = ();
type Rows<'a> = &'a mut Vec<i64>;
type Row<'a> = &'a mut i64;
type WriteToken = ();
fn initialize_skipped_rows(rows: &mut Self::Rows<'_>) {
rows.pop();
}
fn storage_dtype(_params: &Self::Params) -> DType {
DType::from(i64::PTYPE)
}
fn with_capacity(rows: usize, _params: &Self::Params) -> VortexResult<Self> {
Ok(Self(vec![0; rows]))
}
fn rows(&mut self) -> Self::Rows<'_> {
&mut self.0
}
unsafe fn row_unchecked<'a>(rows: &'a mut Self::Rows<'_>, index: usize) -> Self::Row<'a> {
&mut rows[index]
}
unsafe fn finish(self) -> VortexResult<ArrayRef> {
Ok(PrimitiveArray::from_iter(self.0).into_array())
}
}
#[test]
fn test_skip_invalid_sink_rechecks_rows_after_initialization() -> VortexResult<()> {
let input = PrimitiveArray::from_iter([10_i64, 20]).into_array();
let args = VecExecutionArgs::new(vec![input], 2);
let Mask::Values(valid) = Mask::from_iter([false, true]) else {
vortex_bail!("the test validity must be partially valid");
};
let mut ctx = array_session().create_execution_ctx();
let result = execute_sink_valid_rows::<(i64,), (), ShrinkingSink, ()>(
&args,
&valid,
&(),
&mut ctx,
|_| (),
|_, (value,), output| {
*output = value;
},
);
let error = match result {
Err(error) => error,
Ok(_) => vortex_bail!("the sink must reject rows changed by its initializer"),
};
assert!(
error
.to_string()
.contains("initialized output sink must address exactly 2 rows, got 1"),
"unexpected error: {error}",
);
Ok(())
}
}