From 57266ffc2ea7e62c7c2c8bdc13428ac81e5e0a25 Mon Sep 17 00:00:00 2001 From: Huaijin Date: Fri, 7 Aug 2026 14:08:33 +0800 Subject: [PATCH 1/4] feat(array): cast bool and primitive values to Utf8 Signed-off-by: Huaijin --- Cargo.lock | 1 + Cargo.toml | 1 + vortex-array/Cargo.toml | 1 + vortex-array/src/arrays/bool/compute/cast.rs | 111 ++++++++ .../src/arrays/constant/compute/cast.rs | 40 +++ .../src/arrays/primitive/compute/cast.rs | 242 +++++++++++++++++- vortex-array/src/scalar/tests/nested.rs | 43 +++- vortex-array/src/scalar/typed_view/bool.rs | 50 +++- .../src/scalar/typed_view/primitive/scalar.rs | 11 + .../src/scalar/typed_view/primitive/tests.rs | 29 +++ 10 files changed, 508 insertions(+), 21 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 63359686289..df84782675b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -9613,6 +9613,7 @@ dependencies = [ "rstest", "rstest_reuse", "rustc-hash", + "ryu", "serde", "serde_json", "serde_test", diff --git a/Cargo.toml b/Cargo.toml index 36aa5b2ac9e..51c05a25751 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -238,6 +238,7 @@ rstest = "0.26.1" rstest_reuse = "0.7.0" rustc-hash = "2.1.1" rustix = { version = "1.1", features = ["fs"] } +ryu = "1.0.23" serde = "1.0.221" serde_json = "1.0.138" serde_test = "1.0.176" diff --git a/vortex-array/Cargo.toml b/vortex-array/Cargo.toml index d00b811a387..11c3a2cfbd7 100644 --- a/vortex-array/Cargo.toml +++ b/vortex-array/Cargo.toml @@ -53,6 +53,7 @@ regex-syntax = { workspace = true } rstest = { workspace = true, optional = true } rstest_reuse = { workspace = true, optional = true } rustc-hash = { workspace = true } +ryu = { workspace = true } serde = { workspace = true, optional = true, features = ["derive", "rc"] } simdutf8 = { workspace = true } smallvec = { workspace = true } diff --git a/vortex-array/src/arrays/bool/compute/cast.rs b/vortex-array/src/arrays/bool/compute/cast.rs index 36849f55a18..6240cc7f089 100644 --- a/vortex-array/src/arrays/bool/compute/cast.rs +++ b/vortex-array/src/arrays/bool/compute/cast.rs @@ -5,6 +5,7 @@ use num_traits::One; use num_traits::Zero; use vortex_buffer::BufferMut; use vortex_error::VortexResult; +use vortex_mask::Mask; use crate::ArrayRef; use crate::ExecutionCtx; @@ -14,6 +15,8 @@ use crate::arrays::Bool; use crate::arrays::BoolArray; use crate::arrays::PrimitiveArray; use crate::arrays::bool::BoolArrayExt; +use crate::builders::ArrayBuilder; +use crate::builders::VarBinViewBuilder; use crate::dtype::DType; use crate::match_each_native_ptype; use crate::scalar_fn::fns::cast::CastKernel; @@ -53,6 +56,36 @@ impl CastKernel for Bool { )); } + if let DType::Utf8(new_nullability) = dtype { + let len = array.len(); + let new_validity = array + .validity()? + .cast_nullability(*new_nullability, len, ctx)?; + let mask = new_validity.execute_mask(len, ctx)?; + let bits = array.to_bit_buffer(); + let mut builder = VarBinViewBuilder::with_capacity(dtype.clone(), len); + + match &mask { + Mask::AllTrue(_) => { + for value in bits.iter() { + builder.append_value(if value { "true" } else { "false" }); + } + } + Mask::AllFalse(_) => builder.append_nulls(len), + Mask::Values(validity) => { + for (value, valid) in bits.iter().zip(validity.bit_buffer().iter()) { + if valid { + builder.append_value(if value { "true" } else { "false" }); + } else { + builder.append_null(); + } + } + } + } + + return Ok(Some(builder.finish_into_varbinview().into_array())); + } + let DType::Primitive(new_ptype, new_nullability) = dtype else { return Ok(None); }; @@ -79,12 +112,15 @@ mod tests { use std::sync::LazyLock; use rstest::rstest; + use vortex_error::VortexResult; use vortex_session::VortexSession; use crate::Canonical; use crate::IntoArray; use crate::VortexSessionExecute; use crate::arrays::BoolArray; + 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; @@ -117,6 +153,81 @@ mod tests { assert!(result.is_err(), "Expected error, got: {result:?}"); } + #[test] + fn cast_bool_to_utf8() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let actual = BoolArray::from_iter([true, false, true]) + .into_array() + .cast(DType::Utf8(Nullability::NonNullable))?; + let expected = VarBinViewArray::from_iter_str(["true", "false", "true"]); + + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } + + #[test] + fn cast_nullable_bool_to_utf8() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let actual = BoolArray::from_iter([Some(true), None, Some(false)]) + .into_array() + .cast(DType::Utf8(Nullability::Nullable))?; + let expected = VarBinViewArray::from_iter_nullable_str([Some("true"), None, Some("false")]); + + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } + + #[test] + fn cast_all_null_bool_to_utf8() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let actual = BoolArray::from_iter([None, None]) + .into_array() + .cast(DType::Utf8(Nullability::Nullable))?; + let expected = VarBinViewArray::from_iter_nullable_str([None::<&str>, None]); + + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } + + #[test] + fn cast_nullable_bool_with_null_to_non_nullable_utf8_fails() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let result = BoolArray::from_iter([Some(true), None]) + .into_array() + .cast(DType::Utf8(Nullability::NonNullable))? + .execute::(&mut ctx); + + assert!(result.is_err(), "Expected error, got: {result:?}"); + Ok(()) + } + + #[test] + fn cast_all_valid_nullable_bool_to_non_nullable_utf8() -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let actual = BoolArray::from_iter([Some(true), Some(false)]) + .into_array() + .cast(DType::Utf8(Nullability::NonNullable))?; + let expected = VarBinViewArray::from_iter_str(["true", "false"]); + + assert_arrays_eq!(actual, expected, &mut ctx); + Ok(()) + } + + #[test] + fn cast_bool_to_binary_is_unsupported() { + let mut ctx = SESSION.create_execution_ctx(); + let result = BoolArray::from_iter([true, false]) + .into_array() + .cast(DType::Binary(Nullability::NonNullable)) + .and_then(|array| { + array + .execute::(&mut ctx) + .map(|canonical| canonical.into_array()) + }); + + assert!(result.is_err(), "Expected error, got: {result:?}"); + } + #[rstest] #[case(BoolArray::from_iter(vec![true, false, true, true, false]))] #[case(BoolArray::from_iter(vec![Some(true), Some(false), None, Some(true), None]))] diff --git a/vortex-array/src/arrays/constant/compute/cast.rs b/vortex-array/src/arrays/constant/compute/cast.rs index 439bf8367b8..188752491eb 100644 --- a/vortex-array/src/arrays/constant/compute/cast.rs +++ b/vortex-array/src/arrays/constant/compute/cast.rs @@ -23,6 +23,7 @@ impl CastReduce for Constant { #[cfg(test)] mod tests { use rstest::rstest; + use vortex_error::VortexResult; use crate::IntoArray; use crate::VortexSessionExecute; @@ -33,6 +34,7 @@ mod tests { use crate::dtype::DType; use crate::dtype::DecimalDType; use crate::dtype::Nullability; + use crate::dtype::PType; use crate::scalar::DecimalValue; use crate::scalar::Scalar; @@ -65,4 +67,42 @@ mod tests { Some(DecimalValue::I128(4200)) ); } + + #[rstest] + #[case( + Scalar::from(true), + DType::Primitive(PType::I32, Nullability::NonNullable), + Scalar::primitive(1i32, Nullability::NonNullable) + )] + #[case( + Scalar::from(false), + DType::Utf8(Nullability::Nullable), + Scalar::utf8("false", Nullability::Nullable) + )] + #[case( + Scalar::from(-42i64), + DType::Utf8(Nullability::NonNullable), + Scalar::utf8("-42", Nullability::NonNullable) + )] + #[case( + Scalar::from(100.0f64), + DType::Utf8(Nullability::NonNullable), + Scalar::utf8("100.0", Nullability::NonNullable) + )] + #[case( + Scalar::null(DType::Primitive(PType::I64, Nullability::Nullable)), + DType::Utf8(Nullability::Nullable), + Scalar::null(DType::Utf8(Nullability::Nullable)) + )] + fn test_cast_bool_and_primitive_constants( + #[case] source: Scalar, + #[case] target: DType, + #[case] expected: Scalar, + ) -> VortexResult<()> { + let casted = ConstantArray::new(source, 5).into_array().cast(target)?; + + assert_eq!(casted.len(), 5); + assert_eq!(casted.as_constant(), Some(expected)); + Ok(()) + } } diff --git a/vortex-array/src/arrays/primitive/compute/cast.rs b/vortex-array/src/arrays/primitive/compute/cast.rs index 8f5a0f95fcd..d648386115a 100644 --- a/vortex-array/src/arrays/primitive/compute/cast.rs +++ b/vortex-array/src/arrays/primitive/compute/cast.rs @@ -1,6 +1,8 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors +use std::fmt::Write; + use num_traits::AsPrimitive; use num_traits::CheckedMul; use num_traits::NumCast; @@ -24,6 +26,8 @@ use crate::arrays::DecimalArray; use crate::arrays::Primitive; use crate::arrays::PrimitiveArray; use crate::arrays::primitive::PrimitiveArrayExt; +use crate::builders::ArrayBuilder; +use crate::builders::VarBinViewBuilder; use crate::dtype::BigCast; use crate::dtype::DType; use crate::dtype::DecimalDType; @@ -85,10 +89,11 @@ impl CastKernel for Primitive { if let DType::Decimal(decimal_dtype, nullability) = dtype { return cast_to_decimal(array, *decimal_dtype, *nullability, ctx).map(Some); } - let DType::Primitive(new_ptype, new_nullability) = dtype else { - return Ok(None); + let (new_ptype, new_nullability) = match dtype { + DType::Primitive(new_ptype, new_nullability) => (*new_ptype, *new_nullability), + DType::Utf8(_) => return Ok(Some(cast_primitive_to_utf8(array, dtype, ctx)?)), + _ => return Ok(None), }; - let (new_ptype, new_nullability) = (*new_ptype, *new_nullability); let src_ptype = array.ptype(); let new_validity = array @@ -616,13 +621,106 @@ fn cached_values_fit_in(array: ArrayView<'_, Primitive>, target_dtype: &DType) - Some(min.cast(target_dtype).is_ok() && max.cast(target_dtype).is_ok()) } +fn cast_primitive_to_utf8( + array: ArrayView<'_, Primitive>, + dtype: &DType, + ctx: &mut ExecutionCtx, +) -> VortexResult { + let len = array.len(); + let new_validity = array + .validity()? + .cast_nullability(dtype.nullability(), len, ctx)?; + let mask = new_validity.execute_mask(len, ctx)?; + let mut builder = VarBinViewBuilder::with_capacity(dtype.clone(), len); + + // Arrow uses ryu for f32/f64 and Display for f16 and integer values. + match array.ptype() { + PType::F32 => append_float_values_to_utf8(&mut builder, array.as_slice::(), &mask), + PType::F64 => append_float_values_to_utf8(&mut builder, array.as_slice::(), &mask), + ptype => match_each_native_ptype!(ptype, |T| { + append_primitive_values_to_utf8(&mut builder, array.as_slice::(), &mask)?; + }), + } + + Ok(builder.finish_into_varbinview().into_array()) +} + +fn append_float_values_to_utf8( + builder: &mut VarBinViewBuilder, + values: &[T], + mask: &Mask, +) { + let mut formatter = ryu::Buffer::new(); + + match mask { + Mask::AllTrue(_) => { + for &value in values { + builder.append_value(formatter.format(value)); + } + } + Mask::AllFalse(_) => builder.append_nulls(values.len()), + Mask::Values(validity) => { + for (&value, valid) in values.iter().zip(validity.bit_buffer().iter()) { + if valid { + builder.append_value(formatter.format(value)); + } else { + builder.append_null(); + } + } + } + } +} + +fn append_primitive_values_to_utf8( + builder: &mut VarBinViewBuilder, + values: &[T], + mask: &Mask, +) -> VortexResult<()> { + let mut scratch = String::with_capacity(32); + + match mask { + Mask::AllTrue(_) => { + for &value in values { + append_primitive_value(builder, &mut scratch, value)?; + } + } + Mask::AllFalse(_) => builder.append_nulls(values.len()), + Mask::Values(validity) => { + for (&value, valid) in values.iter().zip(validity.bit_buffer().iter()) { + if valid { + append_primitive_value(builder, &mut scratch, value)?; + } else { + builder.append_null(); + } + } + } + } + + Ok(()) +} + +fn append_primitive_value( + builder: &mut VarBinViewBuilder, + scratch: &mut String, + value: T, +) -> VortexResult<()> { + scratch.clear(); + write!(scratch, "{value}") + .map_err(|_| vortex_err!("Failed to format {} value as Utf8", T::PTYPE))?; + builder.append_value(scratch.as_str()); + Ok(()) +} + #[cfg(test)] mod test { + use num_traits::NumCast; use rstest::rstest; use vortex_buffer::BitBuffer; use vortex_buffer::buffer; use vortex_error::VortexError; use vortex_error::VortexResult; + use vortex_error::VortexResult as TestResult; + use vortex_error::vortex_err; use vortex_mask::Mask; use crate::ArrayRef; @@ -631,6 +729,7 @@ mod test { use crate::array_session; use crate::arrays::DecimalArray; use crate::arrays::PrimitiveArray; + use crate::arrays::VarBinViewArray; use crate::assert_arrays_eq; use crate::builtins::ArrayBuiltins; use crate::compute::conformance::cast::test_cast_conformance; @@ -641,6 +740,7 @@ mod test { use crate::dtype::PType; use crate::dtype::i256; use crate::expr::stats::Stat; + use crate::match_each_native_ptype; use crate::validity::Validity; #[test] @@ -1155,4 +1255,140 @@ mod test { fn test_cast_primitive_conformance(#[case] array: ArrayRef) { test_cast_conformance(&array, &mut array_session().create_execution_ctx()); } + + #[rstest] + #[case(PType::U8)] + #[case(PType::U16)] + #[case(PType::U32)] + #[case(PType::U64)] + #[case(PType::I8)] + #[case(PType::I16)] + #[case(PType::I32)] + #[case(PType::I64)] + #[case(PType::F16)] + #[case(PType::F32)] + #[case(PType::F64)] + fn cast_each_primitive_type_to_utf8(#[case] ptype: PType) -> TestResult<()> { + let array = match_each_native_ptype!(ptype, |T| { + let zero = ::from(0u8) + .ok_or_else(|| vortex_err!("Cannot construct zero as {ptype}"))?; + let one = ::from(1u8) + .ok_or_else(|| vortex_err!("Cannot construct one as {ptype}"))?; + let answer = ::from(42u8) + .ok_or_else(|| vortex_err!("Cannot construct 42 as {ptype}"))?; + PrimitiveArray::from_iter([zero, one, answer]).into_array() + }); + let actual = array.cast(DType::Utf8(Nullability::NonNullable))?; + let expected = if matches!(ptype, PType::F32 | PType::F64) { + VarBinViewArray::from_iter_str(["0.0", "1.0", "42.0"]) + } else { + VarBinViewArray::from_iter_str(["0", "1", "42"]) + }; + + assert_arrays_eq!( + actual, + expected, + &mut array_session().create_execution_ctx() + ); + Ok(()) + } + + #[test] + fn cast_nullable_primitive_to_utf8() -> TestResult<()> { + let actual = PrimitiveArray::from_option_iter([Some(100i64), None, Some(-42)]) + .into_array() + .cast(DType::Utf8(Nullability::Nullable))?; + let expected = VarBinViewArray::from_iter_nullable_str([Some("100"), None, Some("-42")]); + + assert_arrays_eq!( + actual, + expected, + &mut array_session().create_execution_ctx() + ); + Ok(()) + } + + #[test] + fn cast_all_null_primitive_to_utf8() -> TestResult<()> { + let actual = PrimitiveArray::from_option_iter([None::, None]) + .into_array() + .cast(DType::Utf8(Nullability::Nullable))?; + let expected = VarBinViewArray::from_iter_nullable_str([None::<&str>, None]); + + assert_arrays_eq!( + actual, + expected, + &mut array_session().create_execution_ctx() + ); + Ok(()) + } + + #[test] + fn cast_nullable_primitive_with_null_to_non_nullable_utf8_fails() -> TestResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let result = PrimitiveArray::from_option_iter([Some(1i64), None]) + .into_array() + .cast(DType::Utf8(Nullability::NonNullable))? + .execute::(&mut ctx); + + assert!(result.is_err(), "Expected error, got: {result:?}"); + Ok(()) + } + + #[test] + fn cast_all_valid_nullable_primitive_to_non_nullable_utf8() -> TestResult<()> { + let actual = PrimitiveArray::from_option_iter([Some(1i64), Some(-42)]) + .into_array() + .cast(DType::Utf8(Nullability::NonNullable))?; + let expected = VarBinViewArray::from_iter_str(["1", "-42"]); + + assert_arrays_eq!( + actual, + expected, + &mut array_session().create_execution_ctx() + ); + Ok(()) + } + + #[test] + fn cast_f64_to_utf8_matches_arrow_formatting() -> TestResult<()> { + let actual = buffer![ + 0.0f64, + -0.0, + 1.5, + 100.0, + 1e20, + 1e-20, + f64::NAN, + f64::INFINITY, + f64::NEG_INFINITY + ] + .into_array() + .cast(DType::Utf8(Nullability::NonNullable))?; + let expected = VarBinViewArray::from_iter_str([ + "0.0", "-0.0", "1.5", "100.0", "1e20", "1e-20", "NaN", "inf", "-inf", + ]); + + assert_arrays_eq!( + actual, + expected, + &mut array_session().create_execution_ctx() + ); + Ok(()) + } + + #[test] + fn cast_primitive_to_binary_is_unsupported() { + let mut ctx = array_session().create_execution_ctx(); + let result = buffer![1i64, 2, 3] + .into_array() + .cast(DType::Binary(Nullability::NonNullable)) + .and_then(|array| { + array + .execute::(&mut ctx) + .map(|canonical| canonical.into_array()) + }); + + assert!(result.is_err(), "Expected error, got: {result:?}"); + } } diff --git a/vortex-array/src/scalar/tests/nested.rs b/vortex-array/src/scalar/tests/nested.rs index d02bf43c631..ad1c9182344 100644 --- a/vortex-array/src/scalar/tests/nested.rs +++ b/vortex-array/src/scalar/tests/nested.rs @@ -7,6 +7,8 @@ mod tests { use std::sync::Arc; + use vortex_error::VortexResult; + use crate::dtype::DType; use crate::dtype::Nullability; use crate::dtype::PType; @@ -505,7 +507,7 @@ mod tests { } #[test] - fn test_list_cast_incompatible_element_types() { + fn test_list_cast_bool_and_primitive_elements_to_utf8() -> VortexResult<()> { // Create a list of integers. let int_list = Scalar::list( Arc::from(DType::Primitive(PType::I32, Nullability::NonNullable)), @@ -513,12 +515,43 @@ mod tests { Nullability::NonNullable, ); - // Try to cast to list of strings - should fail. - let target = DType::List( - Arc::from(DType::Utf8(Nullability::NonNullable)), + let utf8_dtype = DType::Utf8(Nullability::NonNullable); + let target = DType::List(Arc::from(utf8_dtype.clone()), Nullability::NonNullable); + let casted = int_list.cast(&target)?; + let expected = Scalar::list( + utf8_dtype, + vec![Scalar::utf8("1", Nullability::NonNullable)], + Nullability::NonNullable, + ); + + assert_eq!(casted, expected); + + let bool_list = Scalar::list( + Arc::from(DType::Bool(Nullability::NonNullable)), + vec![ + Scalar::bool(true, Nullability::NonNullable), + Scalar::bool(false, Nullability::NonNullable), + ], + Nullability::NonNullable, + ); + let casted = bool_list.cast(&target)?; + let expected = Scalar::list( + DType::Utf8(Nullability::NonNullable), + vec![ + Scalar::utf8("true", Nullability::NonNullable), + Scalar::utf8("false", Nullability::NonNullable), + ], + Nullability::NonNullable, + ); + + assert_eq!(casted, expected); + + let unsupported_target = DType::List( + Arc::from(DType::Bool(Nullability::NonNullable)), Nullability::NonNullable, ); - assert!(int_list.cast(&target).is_err()); + assert!(int_list.cast(&unsupported_target).is_err()); + Ok(()) } #[test] diff --git a/vortex-array/src/scalar/typed_view/bool.rs b/vortex-array/src/scalar/typed_view/bool.rs index c971d1169c6..e9f7e82a068 100644 --- a/vortex-array/src/scalar/typed_view/bool.rs +++ b/vortex-array/src/scalar/typed_view/bool.rs @@ -7,11 +7,14 @@ use std::cmp::Ordering; use std::fmt::Display; use std::fmt::Formatter; +use num_traits::One; +use num_traits::Zero; use vortex_error::VortexExpect; use vortex_error::VortexResult; use vortex_error::vortex_bail; use crate::dtype::DType; +use crate::match_each_native_ptype; use crate::scalar::Scalar; use crate::scalar::ScalarValue; @@ -83,15 +86,19 @@ impl<'a> BoolScalar<'a> { /// Casts this scalar to the given `dtype`. pub(crate) fn cast(&self, dtype: &DType) -> VortexResult { - if !matches!(dtype, DType::Bool(..)) { - vortex_bail!( - "Cannot cast bool to {dtype}: boolean scalars can only be cast to boolean types with different nullability" - ) + let value = self.value.vortex_expect("nullness handled in Scalar::cast"); + + match dtype { + DType::Bool(nullability) => Ok(Scalar::bool(value, *nullability)), + DType::Primitive(ptype, nullability) => Ok(match_each_native_ptype!(*ptype, |T| { + Scalar::primitive(if value { T::one() } else { T::zero() }, *nullability) + })), + DType::Utf8(nullability) => Ok(Scalar::utf8( + if value { "true" } else { "false" }, + *nullability, + )), + _ => vortex_bail!("Cannot cast bool scalar to {dtype}"), } - Ok(Scalar::bool( - self.value.vortex_expect("nullness handled in Scalar::cast"), - dtype.nullability(), - )) } /// Returns a new boolean scalar with the inverted value. @@ -204,14 +211,31 @@ mod test { } #[test] - fn test_bool_cast_to_non_bool_fails() { + fn test_bool_cast_to_primitive_and_utf8() -> VortexResult<()> { use crate::dtype::PType; - let bool_scalar = Scalar::bool(true, NonNullable); - let bool = bool_scalar.as_bool(); + let true_scalar = Scalar::bool(true, NonNullable); + let false_scalar = Scalar::bool(false, NonNullable); - let result = bool.cast(&DType::Primitive(PType::I32, NonNullable)); - assert!(result.is_err()); + assert_eq!( + true_scalar.cast(&DType::Primitive(PType::I32, NonNullable))?, + Scalar::primitive(1i32, NonNullable) + ); + assert_eq!( + false_scalar.cast(&DType::Primitive(PType::F64, Nullable))?, + Scalar::primitive(0.0f64, Nullable) + ); + assert_eq!( + true_scalar.cast(&DType::Utf8(NonNullable))?, + Scalar::utf8("true", NonNullable) + ); + assert_eq!( + false_scalar.cast(&DType::Utf8(Nullable))?, + Scalar::utf8("false", Nullable) + ); + assert!(true_scalar.cast(&DType::Binary(NonNullable)).is_err()); + + Ok(()) } #[test] diff --git a/vortex-array/src/scalar/typed_view/primitive/scalar.rs b/vortex-array/src/scalar/typed_view/primitive/scalar.rs index 3ca4d337a45..4e9cbb4c6ef 100644 --- a/vortex-array/src/scalar/typed_view/primitive/scalar.rs +++ b/vortex-array/src/scalar/typed_view/primitive/scalar.rs @@ -181,6 +181,17 @@ impl<'a> PrimitiveScalar<'a> { *decimal_dtype, *nullability, )), + DType::Utf8(nullability) => { + // Match Arrow's formatting: ryu for f32/f64, Display for f16 and integers. + let value = match self.ptype { + PType::F32 => ryu::Buffer::new().format(pvalue.cast::()?).to_owned(), + PType::F64 => ryu::Buffer::new().format(pvalue.cast::()?).to_owned(), + ptype => { + match_each_native_ptype!(ptype, |T| { pvalue.cast::()?.to_string() }) + } + }; + Ok(Scalar::utf8(value, *nullability)) + } _ => vortex_bail!("Cannot cast primitive scalar to {dtype}"), } } diff --git a/vortex-array/src/scalar/typed_view/primitive/tests.rs b/vortex-array/src/scalar/typed_view/primitive/tests.rs index b0c0063a3df..c61a2ecb5ff 100644 --- a/vortex-array/src/scalar/typed_view/primitive/tests.rs +++ b/vortex-array/src/scalar/typed_view/primitive/tests.rs @@ -6,6 +6,7 @@ use std::cmp::Ordering; use num_traits::CheckedSub; use rstest::rstest; use vortex_error::VortexExpect; +use vortex_error::VortexResult; use vortex_utils::aliases::hash_set::HashSet; use super::pvalue::CoercePValue; @@ -18,6 +19,7 @@ use crate::dtype::ToBytes; use crate::dtype::half::f16; use crate::scalar::PValue; use crate::scalar::PrimitiveScalar; +use crate::scalar::Scalar; use crate::scalar::ScalarValue; #[test] @@ -165,6 +167,33 @@ fn test_primitive_cast( } } +#[rstest] +#[case(Scalar::primitive(42u8, Nullability::NonNullable), "42")] +#[case(Scalar::primitive(-42i64, Nullability::NonNullable), "-42")] +#[case(Scalar::primitive(f16::from_f32(42.0), Nullability::NonNullable), "42")] +#[case(Scalar::primitive(100.0f32, Nullability::NonNullable), "100.0")] +#[case(Scalar::primitive(-0.0f64, Nullability::NonNullable), "-0.0")] +#[case(Scalar::primitive(f64::NAN, Nullability::NonNullable), "NaN")] +#[case(Scalar::primitive(f64::INFINITY, Nullability::NonNullable), "inf")] +#[case(Scalar::primitive(f64::NEG_INFINITY, Nullability::NonNullable), "-inf")] +fn test_primitive_cast_to_utf8(#[case] scalar: Scalar, #[case] expected: &str) -> VortexResult<()> { + let actual = scalar.cast(&DType::Utf8(Nullability::Nullable))?; + + assert_eq!(actual, Scalar::utf8(expected, Nullability::Nullable)); + Ok(()) +} + +#[test] +fn test_primitive_cast_to_binary_fails() { + let scalar = Scalar::primitive(42i64, Nullability::NonNullable); + + assert!( + scalar + .cast(&DType::Binary(Nullability::NonNullable)) + .is_err() + ); +} + #[test] fn test_as_conversion_success() { let dtype = DType::Primitive(PType::I32, Nullability::NonNullable); From 61d9744f097e4eea0f8d53e0d93f5d4d1b692101 Mon Sep 17 00:00:00 2001 From: Huaijin Date: Fri, 7 Aug 2026 14:41:46 +0800 Subject: [PATCH 2/4] refactor(array): use itoa for integer-to-utf8 cast and drop dead error path Signed-off-by: Huaijin --- Cargo.lock | 1 + Cargo.toml | 1 + vortex-array/Cargo.toml | 1 + .../src/arrays/primitive/compute/cast.rs | 97 +++++++++---------- 4 files changed, 50 insertions(+), 50 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index df84782675b..71f595b7a5a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -9596,6 +9596,7 @@ dependencies = [ "insta", "inventory", "itertools 0.14.0", + "itoa", "jiff", "memchr", "mimalloc", diff --git a/Cargo.toml b/Cargo.toml index 51c05a25751..d44899c4e92 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -176,6 +176,7 @@ indicatif = "0.18.0" insta = "1.43" inventory = "0.3.20" itertools = "0.14.0" +itoa = "1.0.18" jiff = "0.2.28" jni = { version = "0.22.0" } kanal = "0.1.1" diff --git a/vortex-array/Cargo.toml b/vortex-array/Cargo.toml index 11c3a2cfbd7..1245997ae67 100644 --- a/vortex-array/Cargo.toml +++ b/vortex-array/Cargo.toml @@ -36,6 +36,7 @@ half = { workspace = true, features = ["num-traits"] } humansize = { workspace = true } inventory = { workspace = true } itertools = { workspace = true } +itoa = { workspace = true } jiff = { workspace = true } memchr = { workspace = true } num-traits = { workspace = true } diff --git a/vortex-array/src/arrays/primitive/compute/cast.rs b/vortex-array/src/arrays/primitive/compute/cast.rs index d648386115a..095bb5e8efa 100644 --- a/vortex-array/src/arrays/primitive/compute/cast.rs +++ b/vortex-array/src/arrays/primitive/compute/cast.rs @@ -38,6 +38,7 @@ use crate::dtype::NativePType; use crate::dtype::Nullability; use crate::dtype::PType; use crate::dtype::ToI256; +use crate::dtype::half::f16; use crate::dtype::i256; use crate::expr::stats::Stat; use crate::expr::stats::StatsProvider; @@ -633,82 +634,78 @@ fn cast_primitive_to_utf8( let mask = new_validity.execute_mask(len, ctx)?; let mut builder = VarBinViewBuilder::with_capacity(dtype.clone(), len); - // Arrow uses ryu for f32/f64 and Display for f16 and integer values. + // Match Arrow's formatting: ryu for f32/f64, Display for f16, and itoa for integers + // (Arrow's lexical_core produces the same output as itoa/Display for integers). match array.ptype() { - PType::F32 => append_float_values_to_utf8(&mut builder, array.as_slice::(), &mask), - PType::F64 => append_float_values_to_utf8(&mut builder, array.as_slice::(), &mask), - ptype => match_each_native_ptype!(ptype, |T| { - append_primitive_values_to_utf8(&mut builder, array.as_slice::(), &mask)?; + PType::F16 => { + let mut scratch = String::with_capacity(16); + append_values_to_utf8( + &mut builder, + array.as_slice::(), + &mask, + |builder, value| { + scratch.clear(); + // Writing to a String is infallible. + let _ = write!(scratch, "{value}"); + builder.append_value(scratch.as_str()); + }, + ); + } + PType::F32 => { + let mut formatter = ryu::Buffer::new(); + append_values_to_utf8( + &mut builder, + array.as_slice::(), + &mask, + |builder, value| builder.append_value(formatter.format(value)), + ); + } + PType::F64 => { + let mut formatter = ryu::Buffer::new(); + append_values_to_utf8( + &mut builder, + array.as_slice::(), + &mask, + |builder, value| builder.append_value(formatter.format(value)), + ); + } + ptype => match_each_integer_ptype!(ptype, |T| { + let mut formatter = itoa::Buffer::new(); + append_values_to_utf8( + &mut builder, + array.as_slice::(), + &mask, + |builder, value| builder.append_value(formatter.format(value)), + ); }), } Ok(builder.finish_into_varbinview().into_array()) } -fn append_float_values_to_utf8( +fn append_values_to_utf8( builder: &mut VarBinViewBuilder, values: &[T], mask: &Mask, + mut append: impl FnMut(&mut VarBinViewBuilder, T), ) { - let mut formatter = ryu::Buffer::new(); - - match mask { - Mask::AllTrue(_) => { - for &value in values { - builder.append_value(formatter.format(value)); - } - } - Mask::AllFalse(_) => builder.append_nulls(values.len()), - Mask::Values(validity) => { - for (&value, valid) in values.iter().zip(validity.bit_buffer().iter()) { - if valid { - builder.append_value(formatter.format(value)); - } else { - builder.append_null(); - } - } - } - } -} - -fn append_primitive_values_to_utf8( - builder: &mut VarBinViewBuilder, - values: &[T], - mask: &Mask, -) -> VortexResult<()> { - let mut scratch = String::with_capacity(32); - match mask { Mask::AllTrue(_) => { for &value in values { - append_primitive_value(builder, &mut scratch, value)?; + append(builder, value); } } Mask::AllFalse(_) => builder.append_nulls(values.len()), Mask::Values(validity) => { for (&value, valid) in values.iter().zip(validity.bit_buffer().iter()) { if valid { - append_primitive_value(builder, &mut scratch, value)?; + append(builder, value); } else { builder.append_null(); } } } } - - Ok(()) -} - -fn append_primitive_value( - builder: &mut VarBinViewBuilder, - scratch: &mut String, - value: T, -) -> VortexResult<()> { - scratch.clear(); - write!(scratch, "{value}") - .map_err(|_| vortex_err!("Failed to format {} value as Utf8", T::PTYPE))?; - builder.append_value(scratch.as_str()); - Ok(()) } #[cfg(test)] From d56165ea15abf2e7717a9cebe18705af5ae6cc7e Mon Sep 17 00:00:00 2001 From: Huaijin Date: Fri, 7 Aug 2026 15:06:03 +0800 Subject: [PATCH 3/4] test(datafusion): cover string cast projection pushdown Signed-off-by: Huaijin --- vortex-datafusion/src/convert/exprs.rs | 60 +++++++++++++++++++++++++- 1 file changed, 59 insertions(+), 1 deletion(-) diff --git a/vortex-datafusion/src/convert/exprs.rs b/vortex-datafusion/src/convert/exprs.rs index f9c4c65a460..b500663835e 100644 --- a/vortex-datafusion/src/convert/exprs.rs +++ b/vortex-datafusion/src/convert/exprs.rs @@ -728,8 +728,13 @@ mod tests { use arrow_schema::Schema; use arrow_schema::TimeUnit as ArrowTimeUnit; use datafusion::arrow::array::AsArray; + use datafusion::arrow::array::BooleanArray; + use datafusion::arrow::array::Float64Array; + use datafusion::arrow::array::Int64Array; + use datafusion::arrow::array::RecordBatch; use datafusion::arrow::datatypes::Int32Type; use datafusion_common::ScalarValue; + use datafusion_common::assert_batches_eq; use datafusion_common::config::ConfigOptions; use datafusion_expr::Operator as DFOperator; use datafusion_expr::ScalarUDF; @@ -1195,7 +1200,7 @@ mod tests { .show() .await?; - // This fails as it pushes string cast to the scan + // Exercise the fallback path with projection pushdown disabled. ctx.session .sql(r#"select cast(id as string) from 'example.vortex'"#) .await? @@ -1205,6 +1210,59 @@ mod tests { Ok(()) } + #[tokio::test] + async fn test_cast_to_string_with_projection_pushdown() -> anyhow::Result<()> { + let ctx = TestSessionContext::new(true); + let batch = RecordBatch::try_from_iter([ + ( + "bool_col", + Arc::new(BooleanArray::from(vec![Some(true), Some(false), None])) as _, + ), + ( + "int_col", + Arc::new(Int64Array::from(vec![Some(42), Some(-7), None])) as _, + ), + ( + "float_col", + Arc::new(Float64Array::from(vec![Some(1.5), Some(-2.25), None])) as _, + ), + ])?; + ctx.write_arrow_batch("files/cast_to_string.vortex", &batch) + .await?; + let provider = ctx + .table_provider("cast_to_string", "/files/", batch.schema()) + .await?; + ctx.session.register_table("cast_to_string", provider)?; + + let actual = ctx + .session + .sql( + "SELECT \ + CAST(bool_col AS STRING) AS b, \ + CAST(int_col AS STRING) AS i, \ + CAST(float_col AS STRING) AS f \ + FROM cast_to_string", + ) + .await? + .collect() + .await?; + + assert_batches_eq!( + [ + "+-------+----+-------+", + "| b | i | f |", + "+-------+----+-------+", + "| true | 42 | 1.5 |", + "| false | -7 | -2.25 |", + "| | | |", + "+-------+----+-------+", + ], + &actual + ); + + Ok(()) + } + /// A cast whose target is a UUID-tagged `FixedSizeBinary(16)` must resolve /// through the dtype extension registry (UUID is registered on the default /// session) instead of the static, non-plugin-aware `DType::from_arrow`, From d12b54a4aa235043a389bd8b913c1f7ba7137df7 Mon Sep 17 00:00:00 2001 From: Huaijin Date: Fri, 7 Aug 2026 15:23:51 +0800 Subject: [PATCH 4/4] refactor(array): clean up utf8 cast style and add f16 special-value test Drop the TestResult alias and inline crate::Canonical qualifications in the primitive cast tests, move to_bit_buffer out of the all-null bool cast path, and cover f16 NaN/inf formatting in the array-path utf8 cast. Signed-off-by: Huaijin --- vortex-array/src/arrays/bool/compute/cast.rs | 4 +- .../src/arrays/primitive/compute/cast.rs | 47 ++++++++++++++----- 2 files changed, 36 insertions(+), 15 deletions(-) diff --git a/vortex-array/src/arrays/bool/compute/cast.rs b/vortex-array/src/arrays/bool/compute/cast.rs index 6240cc7f089..0c49bb75ced 100644 --- a/vortex-array/src/arrays/bool/compute/cast.rs +++ b/vortex-array/src/arrays/bool/compute/cast.rs @@ -62,17 +62,17 @@ impl CastKernel for Bool { .validity()? .cast_nullability(*new_nullability, len, ctx)?; let mask = new_validity.execute_mask(len, ctx)?; - let bits = array.to_bit_buffer(); let mut builder = VarBinViewBuilder::with_capacity(dtype.clone(), len); match &mask { Mask::AllTrue(_) => { - for value in bits.iter() { + for value in array.to_bit_buffer().iter() { builder.append_value(if value { "true" } else { "false" }); } } Mask::AllFalse(_) => builder.append_nulls(len), Mask::Values(validity) => { + let bits = array.to_bit_buffer(); for (value, valid) in bits.iter().zip(validity.bit_buffer().iter()) { if valid { builder.append_value(if value { "true" } else { "false" }); diff --git a/vortex-array/src/arrays/primitive/compute/cast.rs b/vortex-array/src/arrays/primitive/compute/cast.rs index 095bb5e8efa..412f4210dc3 100644 --- a/vortex-array/src/arrays/primitive/compute/cast.rs +++ b/vortex-array/src/arrays/primitive/compute/cast.rs @@ -710,17 +710,16 @@ fn append_values_to_utf8( #[cfg(test)] mod test { - use num_traits::NumCast; use rstest::rstest; use vortex_buffer::BitBuffer; use vortex_buffer::buffer; use vortex_error::VortexError; use vortex_error::VortexResult; - use vortex_error::VortexResult as TestResult; use vortex_error::vortex_err; use vortex_mask::Mask; use crate::ArrayRef; + use crate::Canonical; use crate::IntoArray; use crate::VortexSessionExecute; use crate::array_session; @@ -735,6 +734,7 @@ mod test { use crate::dtype::DecimalType; use crate::dtype::Nullability; use crate::dtype::PType; + use crate::dtype::half::f16; use crate::dtype::i256; use crate::expr::stats::Stat; use crate::match_each_native_ptype; @@ -1265,13 +1265,13 @@ mod test { #[case(PType::F16)] #[case(PType::F32)] #[case(PType::F64)] - fn cast_each_primitive_type_to_utf8(#[case] ptype: PType) -> TestResult<()> { + fn cast_each_primitive_type_to_utf8(#[case] ptype: PType) -> VortexResult<()> { let array = match_each_native_ptype!(ptype, |T| { - let zero = ::from(0u8) + let zero = ::from(0u8) .ok_or_else(|| vortex_err!("Cannot construct zero as {ptype}"))?; - let one = ::from(1u8) + let one = ::from(1u8) .ok_or_else(|| vortex_err!("Cannot construct one as {ptype}"))?; - let answer = ::from(42u8) + let answer = ::from(42u8) .ok_or_else(|| vortex_err!("Cannot construct 42 as {ptype}"))?; PrimitiveArray::from_iter([zero, one, answer]).into_array() }); @@ -1291,7 +1291,7 @@ mod test { } #[test] - fn cast_nullable_primitive_to_utf8() -> TestResult<()> { + fn cast_nullable_primitive_to_utf8() -> VortexResult<()> { let actual = PrimitiveArray::from_option_iter([Some(100i64), None, Some(-42)]) .into_array() .cast(DType::Utf8(Nullability::Nullable))?; @@ -1306,7 +1306,7 @@ mod test { } #[test] - fn cast_all_null_primitive_to_utf8() -> TestResult<()> { + fn cast_all_null_primitive_to_utf8() -> VortexResult<()> { let actual = PrimitiveArray::from_option_iter([None::, None]) .into_array() .cast(DType::Utf8(Nullability::Nullable))?; @@ -1321,19 +1321,19 @@ mod test { } #[test] - fn cast_nullable_primitive_with_null_to_non_nullable_utf8_fails() -> TestResult<()> { + fn cast_nullable_primitive_with_null_to_non_nullable_utf8_fails() -> VortexResult<()> { let mut ctx = array_session().create_execution_ctx(); let result = PrimitiveArray::from_option_iter([Some(1i64), None]) .into_array() .cast(DType::Utf8(Nullability::NonNullable))? - .execute::(&mut ctx); + .execute::(&mut ctx); assert!(result.is_err(), "Expected error, got: {result:?}"); Ok(()) } #[test] - fn cast_all_valid_nullable_primitive_to_non_nullable_utf8() -> TestResult<()> { + fn cast_all_valid_nullable_primitive_to_non_nullable_utf8() -> VortexResult<()> { let actual = PrimitiveArray::from_option_iter([Some(1i64), Some(-42)]) .into_array() .cast(DType::Utf8(Nullability::NonNullable))?; @@ -1348,7 +1348,7 @@ mod test { } #[test] - fn cast_f64_to_utf8_matches_arrow_formatting() -> TestResult<()> { + fn cast_f64_to_utf8_matches_arrow_formatting() -> VortexResult<()> { let actual = buffer![ 0.0f64, -0.0, @@ -1374,6 +1374,27 @@ mod test { Ok(()) } + #[test] + fn cast_f16_to_utf8_matches_arrow_formatting() -> VortexResult<()> { + let actual = buffer![ + f16::from_f32(0.0), + f16::from_f32(-42.5), + f16::NAN, + f16::INFINITY, + f16::NEG_INFINITY + ] + .into_array() + .cast(DType::Utf8(Nullability::NonNullable))?; + let expected = VarBinViewArray::from_iter_str(["0", "-42.5", "NaN", "inf", "-inf"]); + + assert_arrays_eq!( + actual, + expected, + &mut array_session().create_execution_ctx() + ); + Ok(()) + } + #[test] fn cast_primitive_to_binary_is_unsupported() { let mut ctx = array_session().create_execution_ctx(); @@ -1382,7 +1403,7 @@ mod test { .cast(DType::Binary(Nullability::NonNullable)) .and_then(|array| { array - .execute::(&mut ctx) + .execute::(&mut ctx) .map(|canonical| canonical.into_array()) });