Skip to content
Closed
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
86 changes: 82 additions & 4 deletions vortex-array/src/arrays/varbinview/compute/cast.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
use std::sync::Arc;

use vortex_error::VortexResult;
use vortex_error::vortex_bail;

use crate::ArrayRef;
use crate::ExecutionCtx;
Expand All @@ -12,6 +13,7 @@ use crate::array::ArrayView;
use crate::arrays::VarBinView;
use crate::arrays::VarBinViewArray;
use crate::dtype::DType;
use crate::dtype::Nullability;
use crate::scalar_fn::fns::cast::CastKernel;
use crate::scalar_fn::fns::cast::CastReduce;
use crate::validity::Validity;
Expand All @@ -21,7 +23,9 @@ fn build_with_validity(
new_dtype: DType,
new_validity: Validity,
) -> ArrayRef {
// SAFETY: casting just changes the DType, does not affect invariants on views/buffers.
// SAFETY: views and buffers are unchanged. Null views may have invalid bounds or UTF-8, so
// removing nullability requires every source row to be valid. The caller proves this from
// the validity representation or its executed bits before exposing those views.
unsafe {
VarBinViewArray::new_handle_unchecked(
array.views_handle().clone(),
Expand Down Expand Up @@ -62,9 +66,16 @@ impl CastKernel for VarBinView {
}

let new_nullability = dtype.nullability();
let new_validity = array
.validity()?
.cast_nullability(new_nullability, array.len(), ctx)?;
let validity = array.validity()?;
let new_validity = match new_nullability {
Nullability::NonNullable => {
if !validity.execute_no_nulls(array.len(), ctx)? {
vortex_bail!(InvalidArgument: "Cannot cast array with invalid values to non-nullable type.");
}
Validity::NonNullable
}
Nullability::Nullable => validity.into_nullable(),
};
let new_dtype = array.dtype().with_nullability(new_nullability);
Ok(Some(build_with_validity(array, new_dtype, new_validity)))
}
Expand All @@ -75,19 +86,86 @@ mod tests {
use std::sync::LazyLock;

use rstest::rstest;
use vortex_buffer::buffer;
use vortex_error::VortexResult;
use vortex_session::VortexSession;

use crate::Canonical;
use crate::IntoArray;
use crate::VortexSessionExecute;
use crate::aggregate_fn::fns::min::MIN_SKIP_NANS;
use crate::arrays::BoolArray;
use crate::arrays::VarBinArray;
use crate::arrays::VarBinViewArray;
use crate::assert_arrays_eq;
use crate::builtins::ArrayBuiltins;
use crate::compute::conformance::cast::test_cast_conformance;
use crate::dtype::DType;
use crate::dtype::Nullability;
use crate::validity::Validity;

static SESSION: LazyLock<VortexSession> = LazyLock::new(crate::array_session);

#[rstest]
#[case::ascii_utf8(DType::Utf8(Nullability::Nullable), b'x')]
#[case::ascii_binary(DType::Binary(Nullability::Nullable), b'x')]
#[case::invalid_utf8(DType::Utf8(Nullability::Nullable), 0xff)]
#[case::non_utf8_null_binary(DType::Binary(Nullability::Nullable), 0xff)]
fn non_nullable_cast_checks_validity_bits(
#[case] dtype: DType,
#[case] null_byte: u8,
) -> VortexResult<()> {
let mut ctx = SESSION.create_execution_ctx();
let validity = BoolArray::from_iter([false, true]).into_array();
let array = VarBinArray::try_new(
buffer![0u32, 1, 2].into_array(),
buffer![null_byte, b'a'],
dtype.clone(),
Validity::Array(validity.clone()),
)?
.into_array()
.execute::<VarBinViewArray>(&mut ctx)?;
let donor = BoolArray::from_iter([true, true]).into_array();
donor
.aggregations()
.compute_result(&MIN_SKIP_NANS, &mut ctx)?;
// The donor minimum is correct for a different input, so it cannot prove this validity.
validity.aggregations().inherit_from(donor.aggregations());

let casted = array.into_array().cast(dtype.as_nonnullable())?;
assert!(casted.execute::<Canonical>(&mut ctx).is_err());
Ok(())
}

#[rstest]
#[case(DType::Utf8(Nullability::Nullable))]
#[case(DType::Binary(Nullability::Nullable))]
fn non_nullable_cast_accepts_all_valid_bits(#[case] dtype: DType) -> VortexResult<()> {
let mut ctx = SESSION.create_execution_ctx();
let validity = BoolArray::from_iter([true, true]).into_array();
let array = VarBinArray::try_new(
buffer![0u32, 1, 2].into_array(),
buffer![b'x', b'a'],
dtype.clone(),
Validity::Array(validity.clone()),
)?
.into_array()
.execute::<VarBinViewArray>(&mut ctx)?;
let donor = BoolArray::from_iter([false, true]).into_array();
donor
.aggregations()
.compute_result(&MIN_SKIP_NANS, &mut ctx)?;
// The donor minimum is correct for a different input, so it cannot prove this validity.
validity.aggregations().inherit_from(donor.aggregations());

let target = dtype.as_nonnullable();
let casted = array.into_array().cast(target.clone())?;
let executed = casted.execute::<Canonical>(&mut ctx)?;
let expected = VarBinViewArray::from_iter([Some("x"), Some("a")], target);
assert_arrays_eq!(executed.into_array(), expected, &mut ctx);
Ok(())
}

#[rstest]
#[case(
DType::Utf8(Nullability::Nullable),
Expand Down
Loading