Skip to main content

vortex_fastlanes/rle/compute/
cast.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use vortex_array::ArrayRef;
5use vortex_array::ArrayView;
6use vortex_array::IntoArray;
7use vortex_array::builtins::ArrayBuiltins;
8use vortex_array::dtype::DType;
9use vortex_array::dtype::Nullability;
10use vortex_array::scalar_fn::fns::cast::CastReduce;
11use vortex_error::VortexResult;
12
13use crate::rle::RLE;
14use crate::rle::RLEArrayExt;
15use crate::rle::RLEArraySlotsExt;
16
17impl CastReduce for RLE {
18    fn cast(array: ArrayView<'_, Self>, dtype: &DType) -> VortexResult<Option<ArrayRef>> {
19        // Cast RLE values.
20        let casted_values = array
21            .values()
22            .cast(DType::Primitive(dtype.as_ptype(), Nullability::NonNullable))?;
23
24        // Cast RLE indices such that validity matches the target dtype.
25        let casted_indices = array.indices().cast(
26            array
27                .indices()
28                .dtype()
29                .with_nullability(dtype.nullability()),
30        )?;
31
32        Ok(Some(
33            RLE::try_new(
34                casted_values,
35                casted_indices,
36                array.values_idx_offsets().clone(),
37                array.offset(),
38                array.len(),
39            )?
40            .into_array(),
41        ))
42    }
43}
44
45#[cfg(test)]
46mod tests {
47    use std::sync::LazyLock;
48
49    use rstest::rstest;
50    use vortex_array::Canonical;
51    use vortex_array::ExecutionCtx;
52    use vortex_array::IntoArray;
53    use vortex_array::VortexSessionExecute;
54    use vortex_array::arrays::PrimitiveArray;
55    use vortex_array::assert_arrays_eq;
56    use vortex_array::builtins::ArrayBuiltins;
57    use vortex_array::compute::conformance::cast::test_cast_conformance;
58    use vortex_array::dtype::DType;
59    use vortex_array::dtype::Nullability;
60    use vortex_array::dtype::PType;
61    use vortex_array::validity::Validity;
62    use vortex_buffer::Buffer;
63    use vortex_session::VortexSession;
64
65    use crate::RLEData;
66    use crate::rle::RLEArray;
67
68    static SESSION: LazyLock<VortexSession> = LazyLock::new(|| {
69        let session = vortex_array::array_session();
70        crate::initialize(&session);
71        session
72    });
73
74    fn rle(primitive: &PrimitiveArray, ctx: &mut ExecutionCtx) -> RLEArray {
75        RLEData::encode(primitive.as_view(), ctx).unwrap()
76    }
77
78    #[test]
79    fn try_cast_rle_success() {
80        let mut ctx = SESSION.create_execution_ctx();
81        let primitive = PrimitiveArray::new(
82            Buffer::from_iter([10u8, 20, 30, 40, 50]),
83            Validity::from_iter([true, true, true, true, true]),
84        );
85        let encoded = rle(&primitive, &mut ctx);
86
87        let casted = encoded
88            .into_array()
89            .cast(DType::Primitive(PType::U16, Nullability::NonNullable))
90            .unwrap();
91        assert_arrays_eq!(
92            casted,
93            PrimitiveArray::from_iter([10u16, 20, 30, 40, 50]),
94            &mut ctx
95        );
96    }
97
98    #[test]
99    #[should_panic]
100    fn try_cast_rle_fail() {
101        let mut ctx = SESSION.create_execution_ctx();
102        let primitive = PrimitiveArray::new(
103            Buffer::from_iter([10u8, 20, 30, 40, 50]),
104            Validity::from_iter([true, false, true, true, false]),
105        );
106        let encoded = rle(&primitive, &mut ctx);
107        let result = encoded
108            .into_array()
109            .cast(DType::Primitive(PType::U8, Nullability::NonNullable))
110            .and_then(|a| a.execute::<Canonical>(&mut ctx).map(|c| c.into_array()));
111        result.unwrap();
112    }
113
114    #[rstest]
115    #[case::u8(
116        PrimitiveArray::new(
117            Buffer::from_iter([0u8, 10, 20, 30, 40, 50]),
118            Validity::NonNullable,
119        )
120    )]
121    #[case::u8_nullable(
122        PrimitiveArray::new(
123            Buffer::from_iter([0u8, 10, 20, 30, 40]),
124            Validity::from_iter([true, false, true, false, true]),
125        )
126    )]
127    #[case::u16(
128        PrimitiveArray::new(
129            Buffer::from_iter([0u16, 100, 200, 300, 400, 500]),
130            Validity::NonNullable,
131        )
132    )]
133    #[case::u16_nullable(
134        PrimitiveArray::new(
135            Buffer::from_iter([0u16, 100, 200, 300, 400]),
136            Validity::from_iter([false, true, false, true, true]),
137        )
138    )]
139    #[case::u32(
140        PrimitiveArray::new(
141            Buffer::from_iter([0u32, 1000, 2000, 3000, 4000]),
142            Validity::NonNullable,
143        )
144    )]
145    #[case::u32_nullable(
146        PrimitiveArray::new(
147            Buffer::from_iter([0u32, 1000, 2000, 3000, 4000]),
148            Validity::from_iter([true, true, false, false, true]),
149        )
150    )]
151    #[case::u64(
152        PrimitiveArray::new(
153            Buffer::from_iter([0u64, 10000, 20000, 30000]),
154            Validity::NonNullable,
155        )
156    )]
157    #[case::u64_nullable(
158        PrimitiveArray::new(
159            Buffer::from_iter([0u64, 10000, 20000, 30000]),
160            Validity::from_iter([false, false, true, true]),
161        )
162    )]
163    fn test_cast_rle_conformance(#[case] primitive: PrimitiveArray) {
164        let mut ctx = SESSION.create_execution_ctx();
165        let rle_array = rle(&primitive, &mut ctx);
166        test_cast_conformance(&rle_array.into_array(), &mut ctx);
167    }
168}