Skip to main content

vortex_fsst/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::ExecutionCtx;
7use vortex_array::IntoArray;
8use vortex_array::arrays::VarBinArray;
9use vortex_array::arrays::varbin::VarBinArraySlotsExt;
10use vortex_array::dtype::DType;
11use vortex_array::scalar_fn::fns::cast::CastKernel;
12use vortex_array::scalar_fn::fns::cast::CastReduce;
13use vortex_array::validity::Validity;
14use vortex_error::VortexResult;
15
16use crate::FSST;
17use crate::FSSTArrayExt;
18use crate::FSSTArraySlotsExt;
19
20fn build_with_codes_validity(
21    array: ArrayView<'_, FSST>,
22    dtype: &DType,
23    new_codes_validity: Validity,
24) -> VortexResult<ArrayRef> {
25    let codes = array.codes();
26    let new_codes = VarBinArray::try_new(
27        codes.offsets().clone(),
28        codes.bytes().clone(),
29        codes.dtype().with_nullability(dtype.nullability()),
30        new_codes_validity,
31    )?;
32
33    Ok(unsafe {
34        FSST::new_unchecked_with_symbol_table(
35            dtype.clone(),
36            array.symbol_table(),
37            new_codes,
38            array.uncompressed_lengths().clone(),
39        )
40    }
41    .into_array())
42}
43
44impl CastReduce for FSST {
45    fn cast(array: ArrayView<'_, Self>, dtype: &DType) -> VortexResult<Option<ArrayRef>> {
46        if !array.dtype().eq_ignore_nullability(dtype) {
47            return Ok(None);
48        }
49
50        let codes = array.codes();
51        let Some(new_codes_validity) = codes
52            .validity()?
53            .trivially_cast_nullability(dtype.nullability(), codes.len())?
54        else {
55            return Ok(None);
56        };
57
58        Ok(Some(build_with_codes_validity(
59            array,
60            dtype,
61            new_codes_validity,
62        )?))
63    }
64}
65
66impl CastKernel for FSST {
67    fn cast(
68        array: ArrayView<'_, Self>,
69        dtype: &DType,
70        ctx: &mut ExecutionCtx,
71    ) -> VortexResult<Option<ArrayRef>> {
72        if !array.dtype().eq_ignore_nullability(dtype) {
73            return Ok(None);
74        }
75
76        let codes = array.codes();
77        let new_codes_validity =
78            codes
79                .validity()?
80                .cast_nullability(dtype.nullability(), codes.len(), ctx)?;
81
82        Ok(Some(build_with_codes_validity(
83            array,
84            dtype,
85            new_codes_validity,
86        )?))
87    }
88}
89
90#[cfg(test)]
91mod tests {
92    use std::sync::LazyLock;
93
94    use rstest::rstest;
95    use vortex_array::IntoArray;
96    use vortex_array::VortexSessionExecute;
97    use vortex_array::arrays::VarBinArray;
98    use vortex_array::builtins::ArrayBuiltins;
99    use vortex_array::compute::conformance::cast::test_cast_conformance;
100    use vortex_array::dtype::DType;
101    use vortex_array::dtype::Nullability;
102    use vortex_error::VortexResult;
103    use vortex_session::VortexSession;
104
105    use crate::fsst_compress;
106    use crate::fsst_train_compressor;
107    use crate::initialize;
108
109    static SESSION: LazyLock<VortexSession> = LazyLock::new(|| {
110        let session = vortex_array::array_session();
111        initialize(&session);
112        session
113    });
114
115    #[test]
116    fn test_cast_fsst_nullability() -> VortexResult<()> {
117        let mut ctx = SESSION.create_execution_ctx();
118        let strings = VarBinArray::from_iter(
119            vec![Some("hello"), Some("world"), Some("hello world")],
120            DType::Utf8(Nullability::NonNullable),
121        )
122        .into_array();
123
124        let compressor = fsst_train_compressor(&strings, &mut ctx)?;
125        let fsst = fsst_compress(&strings, &compressor, &mut ctx)?;
126
127        // Cast to nullable
128        let casted = fsst.into_array().cast(DType::Utf8(Nullability::Nullable))?;
129        assert_eq!(casted.dtype(), &DType::Utf8(Nullability::Nullable));
130        Ok(())
131    }
132
133    #[rstest]
134    #[case(VarBinArray::from_iter(
135        vec![Some("hello"), Some("world"), Some("hello world")],
136        DType::Utf8(Nullability::NonNullable)
137    ))]
138    #[case(VarBinArray::from_iter(
139        vec![Some("foo"), None, Some("bar"), Some("foobar")],
140        DType::Utf8(Nullability::Nullable)
141    ))]
142    #[case(VarBinArray::from_iter(
143        vec![Some("test")],
144        DType::Utf8(Nullability::NonNullable)
145    ))]
146    fn test_cast_fsst_conformance(#[case] array: VarBinArray) -> VortexResult<()> {
147        let mut ctx = SESSION.create_execution_ctx();
148        let array = array.into_array();
149        let compressor = fsst_train_compressor(&array, &mut ctx)?;
150        let fsst = fsst_compress(&array, &compressor, &mut ctx)?;
151        test_cast_conformance(&fsst.into_array(), &mut ctx);
152        Ok(())
153    }
154}