use vortex_error::VortexResult;
use vortex_error::vortex_ensure;
use vortex_error::vortex_ensure_eq;
use super::super::RowFnExecutionArgs;
use crate::ArrayRef;
use crate::ExecutionCtx;
use crate::IntoArray;
use crate::arrays::ConstantArray;
use crate::builtins::ArrayBuiltins;
use crate::dtype::DType;
use crate::scalar::Scalar;
use crate::scalar_fn::ScalarFnId;
impl RowFnExecutionArgs {
pub(super) fn all_null(&self) -> ArrayRef {
ConstantArray::new(Scalar::null(self.result_dtype.clone()), self.row_count).into_array()
}
pub(super) fn finalize_output(
&self,
values: ArrayRef,
expected_len: usize,
) -> VortexResult<ArrayRef> {
validate_output(self.id, &self.result_dtype, expected_len, &values)?;
cast_output_nullability(&self.result_dtype, values)
}
pub(super) fn validate_kernel_output(
&self,
values: ArrayRef,
expected_len: usize,
ctx: &mut ExecutionCtx,
) -> VortexResult<ArrayRef> {
finalize_kernel_output(self.id, &self.output_dtype, expected_len, values, ctx)
}
}
pub(crate) fn finalize_kernel_output(
id: ScalarFnId,
result_dtype: &DType,
expected_len: usize,
values: ArrayRef,
ctx: &mut ExecutionCtx,
) -> VortexResult<ArrayRef> {
validate_output(id, result_dtype, expected_len, &values)?;
vortex_ensure!(
values.all_valid(ctx)?,
"the {id} row kernel must produce only valid rows, got at least one null row",
);
cast_output_nullability(result_dtype, values)
}
fn validate_output(
id: ScalarFnId,
result_dtype: &DType,
expected_len: usize,
values: &ArrayRef,
) -> VortexResult<()> {
vortex_ensure_eq!(
values.len(),
expected_len,
"the {id} kernel output must contain {expected_len} rows, got {}",
values.len(),
);
let values_with_result_nullability =
values.dtype().with_nullability(result_dtype.nullability());
vortex_ensure!(
values_with_result_nullability == *result_dtype,
"the {id} output dtype must match {result_dtype} except for outer nullability, got {}",
values.dtype(),
);
Ok(())
}
fn cast_output_nullability(result_dtype: &DType, values: ArrayRef) -> VortexResult<ArrayRef> {
if values.dtype() == result_dtype {
Ok(values)
} else {
values.cast(result_dtype.clone())
}
}