|
| 1 | +// SPDX-License-Identifier: Apache-2.0 |
| 2 | +// SPDX-FileCopyrightText: Copyright the Vortex contributors |
| 3 | + |
| 4 | +//! Serialization plugin for the frozen bit-packed wire format. |
| 5 | +
|
| 6 | +use prost::Message; |
| 7 | +use vortex_array::Array; |
| 8 | +use vortex_array::ArrayDeserialization; |
| 9 | +use vortex_array::ArrayId; |
| 10 | +use vortex_array::ArrayParts; |
| 11 | +use vortex_array::ArrayPlugin; |
| 12 | +use vortex_array::ArrayRef; |
| 13 | +use vortex_array::ArraySerialization; |
| 14 | +use vortex_array::ArraySlots; |
| 15 | +use vortex_array::ArrayVTable; |
| 16 | +use vortex_array::IntoArray; |
| 17 | +use vortex_array::patches::Patches; |
| 18 | +use vortex_array::patches::PatchesData; |
| 19 | +use vortex_array::patches::PatchesMetadata; |
| 20 | +use vortex_array::validity::Validity; |
| 21 | +use vortex_array::vtable::validity_to_child; |
| 22 | +use vortex_error::VortexResult; |
| 23 | +use vortex_error::vortex_bail; |
| 24 | +use vortex_error::vortex_ensure; |
| 25 | +use vortex_error::vortex_err; |
| 26 | +use vortex_session::VortexSession; |
| 27 | + |
| 28 | +use crate::BitPacked; |
| 29 | +use crate::BitPackedArrayExt; |
| 30 | +use crate::BitPackedData; |
| 31 | + |
| 32 | +/// Metadata of the frozen `fastlanes.bitpacked` wire format. |
| 33 | +#[derive(Clone, prost::Message)] |
| 34 | +pub(crate) struct BitPackedMetadata { |
| 35 | + #[prost(uint32, tag = "1")] |
| 36 | + pub(crate) bit_width: u32, |
| 37 | + #[prost(uint32, tag = "2")] |
| 38 | + pub(crate) offset: u32, // must be <1024 |
| 39 | + #[prost(message, optional, tag = "3")] |
| 40 | + pub(crate) patches: Option<PatchesMetadata>, |
| 41 | +} |
| 42 | + |
| 43 | +/// Serialization boundary for the frozen `fastlanes.bitpacked` wire format. |
| 44 | +#[derive(Debug, Clone)] |
| 45 | +pub struct BitPackedPlugin; |
| 46 | + |
| 47 | +impl ArrayPlugin for BitPackedPlugin { |
| 48 | + fn id(&self) -> ArrayId { |
| 49 | + ArrayVTable::id(&BitPacked) |
| 50 | + } |
| 51 | + |
| 52 | + fn serialize( |
| 53 | + &self, |
| 54 | + array: &ArrayRef, |
| 55 | + _session: &VortexSession, |
| 56 | + ) -> VortexResult<Option<ArraySerialization>> { |
| 57 | + vortex_ensure!( |
| 58 | + self.id() == array.encoding_id(), |
| 59 | + "array plugin {} cannot serialize in-memory array {}", |
| 60 | + self.id(), |
| 61 | + array.encoding_id(), |
| 62 | + ); |
| 63 | + let view = array.as_::<BitPacked>(); |
| 64 | + let metadata = BitPackedMetadata { |
| 65 | + bit_width: view.bit_width() as u32, |
| 66 | + offset: view.offset() as u32, |
| 67 | + patches: view |
| 68 | + .patches() |
| 69 | + .map(|p| p.to_metadata(view.len(), view.dtype())) |
| 70 | + .transpose()?, |
| 71 | + } |
| 72 | + .encode_to_vec(); |
| 73 | + Ok(Some(ArraySerialization::from_array( |
| 74 | + self.id(), |
| 75 | + array, |
| 76 | + metadata, |
| 77 | + ))) |
| 78 | + } |
| 79 | + |
| 80 | + fn deserialize( |
| 81 | + &self, |
| 82 | + parts: ArrayDeserialization<'_>, |
| 83 | + _session: &VortexSession, |
| 84 | + ) -> VortexResult<ArrayRef> { |
| 85 | + vortex_ensure!( |
| 86 | + self.id() == parts.serialized_id, |
| 87 | + "array plugin {} does not recognize serialized ID {}", |
| 88 | + self.id(), |
| 89 | + parts.serialized_id, |
| 90 | + ); |
| 91 | + let ArrayDeserialization { |
| 92 | + dtype, |
| 93 | + len, |
| 94 | + metadata, |
| 95 | + buffers, |
| 96 | + children, |
| 97 | + .. |
| 98 | + } = parts; |
| 99 | + |
| 100 | + let metadata = BitPackedMetadata::decode(metadata)?; |
| 101 | + if buffers.len() != 1 { |
| 102 | + vortex_bail!("Expected 1 buffer, got {}", buffers.len()); |
| 103 | + } |
| 104 | + let packed = buffers[0].clone(); |
| 105 | + |
| 106 | + let load_validity = |child_idx: usize| { |
| 107 | + if children.len() == child_idx { |
| 108 | + Ok(Validity::from(dtype.nullability())) |
| 109 | + } else if children.len() == child_idx + 1 { |
| 110 | + let validity = children.get(child_idx, &Validity::DTYPE, len)?; |
| 111 | + Ok(Validity::Array(validity)) |
| 112 | + } else { |
| 113 | + vortex_bail!( |
| 114 | + "Expected {} or {} children, got {}", |
| 115 | + child_idx, |
| 116 | + child_idx + 1, |
| 117 | + children.len() |
| 118 | + ); |
| 119 | + } |
| 120 | + }; |
| 121 | + |
| 122 | + let validity_idx = match &metadata.patches { |
| 123 | + None => 0, |
| 124 | + Some(patches_meta) if patches_meta.chunk_offsets_dtype()?.is_some() => 3, |
| 125 | + Some(_) => 2, |
| 126 | + }; |
| 127 | + |
| 128 | + let validity = load_validity(validity_idx)?; |
| 129 | + |
| 130 | + let patches = metadata |
| 131 | + .patches |
| 132 | + .map(|p| { |
| 133 | + let indices = children.get(0, &p.indices_dtype()?, p.len()?)?; |
| 134 | + let values = children.get(1, dtype, p.len()?)?; |
| 135 | + let chunk_offsets = p |
| 136 | + .chunk_offsets_dtype()? |
| 137 | + .map(|dtype| children.get(2, &dtype, p.chunk_offsets_len() as usize)) |
| 138 | + .transpose()?; |
| 139 | + |
| 140 | + Patches::new(len, p.offset()?, indices, values, chunk_offsets) |
| 141 | + }) |
| 142 | + .transpose()?; |
| 143 | + |
| 144 | + let slots = { |
| 145 | + let mut s = ArraySlots::with_capacity(4); |
| 146 | + PatchesData::push_slots(&mut s, patches.as_ref()); |
| 147 | + s.push(validity_to_child(&validity, len)); |
| 148 | + s |
| 149 | + }; |
| 150 | + let data = BitPackedData::try_new( |
| 151 | + packed, |
| 152 | + patches, |
| 153 | + u8::try_from(metadata.bit_width).map_err(|_| { |
| 154 | + vortex_err!( |
| 155 | + "BitPackedMetadata bit_width {} does not fit in u8", |
| 156 | + metadata.bit_width |
| 157 | + ) |
| 158 | + })?, |
| 159 | + u16::try_from(metadata.offset).map_err(|_| { |
| 160 | + vortex_err!( |
| 161 | + "BitPackedMetadata offset {} does not fit in u16", |
| 162 | + metadata.offset |
| 163 | + ) |
| 164 | + })?, |
| 165 | + )?; |
| 166 | + Ok(Array::<BitPacked>::try_from_parts( |
| 167 | + ArrayParts::new(BitPacked, dtype.clone(), len, data).with_slots(slots), |
| 168 | + )? |
| 169 | + .into_array()) |
| 170 | + } |
| 171 | +} |
| 172 | + |
| 173 | +#[cfg(test)] |
| 174 | +mod tests { |
| 175 | + use prost::Message; |
| 176 | + use rstest::rstest; |
| 177 | + use vortex_array::ArrayDeserialization; |
| 178 | + use vortex_array::ArrayPlugin; |
| 179 | + use vortex_array::ArrayVTable; |
| 180 | + use vortex_array::IntoArray; |
| 181 | + use vortex_array::VortexSessionExecute; |
| 182 | + use vortex_array::arrays::PrimitiveArray; |
| 183 | + use vortex_array::assert_arrays_eq; |
| 184 | + use vortex_array::buffer::BufferHandle; |
| 185 | + use vortex_buffer::ByteBuffer; |
| 186 | + use vortex_error::VortexResult; |
| 187 | + use vortex_error::vortex_err; |
| 188 | + |
| 189 | + use super::BitPackedMetadata; |
| 190 | + use super::BitPackedPlugin; |
| 191 | + use crate::BitPacked; |
| 192 | + use crate::BitPackedData; |
| 193 | + |
| 194 | + #[test] |
| 195 | + fn serialize_rejects_other_encodings() { |
| 196 | + let session = vortex_array::array_session(); |
| 197 | + let values = PrimitiveArray::from_iter([1u32, 2, 3]).into_array(); |
| 198 | + assert!(BitPackedPlugin.serialize(&values, &session).is_err()); |
| 199 | + } |
| 200 | + |
| 201 | + #[rstest] |
| 202 | + #[case(2, u32::MAX, u32::MAX, "Expected 0 or 1 children")] |
| 203 | + #[case(0, u32::MAX, u32::MAX, "bit_width")] |
| 204 | + #[case(0, 0, u32::MAX, "offset")] |
| 205 | + fn invalid_inputs_preserve_validation_order( |
| 206 | + #[case] num_children: usize, |
| 207 | + #[case] bit_width: u32, |
| 208 | + #[case] offset: u32, |
| 209 | + #[case] expected: &str, |
| 210 | + ) -> VortexResult<()> { |
| 211 | + let session = vortex_array::array_session(); |
| 212 | + let child = PrimitiveArray::from_iter([0u32]).into_array(); |
| 213 | + let metadata = BitPackedMetadata { |
| 214 | + bit_width, |
| 215 | + offset, |
| 216 | + patches: None, |
| 217 | + } |
| 218 | + .encode_to_vec(); |
| 219 | + let buffers = [BufferHandle::new_host(ByteBuffer::empty())]; |
| 220 | + let children = vec![child.clone(); num_children]; |
| 221 | + let error = BitPackedPlugin |
| 222 | + .deserialize( |
| 223 | + ArrayDeserialization::new( |
| 224 | + BitPackedPlugin.id(), |
| 225 | + child.dtype(), |
| 226 | + 0, |
| 227 | + &metadata, |
| 228 | + &buffers, |
| 229 | + &children, |
| 230 | + ), |
| 231 | + &session, |
| 232 | + ) |
| 233 | + .err() |
| 234 | + .ok_or_else(|| vortex_err!("Invalid inputs must be rejected"))?; |
| 235 | + assert!(error.to_string().contains(expected), "{error}"); |
| 236 | + Ok(()) |
| 237 | + } |
| 238 | + |
| 239 | + #[test] |
| 240 | + fn serde_requires_plugin() -> VortexResult<()> { |
| 241 | + let session = vortex_array::array_session(); |
| 242 | + let mut ctx = session.create_execution_ctx(); |
| 243 | + let values = |
| 244 | + PrimitiveArray::from_option_iter([Some(1u32), None, Some(511), Some(7)]).into_array(); |
| 245 | + let packed = BitPackedData::encode(&values, 3, &mut ctx)?; |
| 246 | + let array = packed.as_array(); |
| 247 | + let serialized = BitPackedPlugin |
| 248 | + .serialize(array, &session)? |
| 249 | + .ok_or_else(|| vortex_err!("BitPacked must serialize"))?; |
| 250 | + let buffers = serialized |
| 251 | + .buffers |
| 252 | + .iter() |
| 253 | + .cloned() |
| 254 | + .map(BufferHandle::new_host) |
| 255 | + .collect::<Vec<_>>(); |
| 256 | + |
| 257 | + assert!(ArrayVTable::serialize(packed.as_view(), &session).is_err()); |
| 258 | + assert!( |
| 259 | + ArrayVTable::deserialize( |
| 260 | + &BitPacked, |
| 261 | + array.dtype(), |
| 262 | + array.len(), |
| 263 | + &serialized.metadata, |
| 264 | + &buffers, |
| 265 | + &serialized.children, |
| 266 | + &session, |
| 267 | + ) |
| 268 | + .is_err() |
| 269 | + ); |
| 270 | + |
| 271 | + let decoded = BitPackedPlugin.deserialize( |
| 272 | + ArrayDeserialization::new( |
| 273 | + BitPackedPlugin.id(), |
| 274 | + array.dtype(), |
| 275 | + array.len(), |
| 276 | + &serialized.metadata, |
| 277 | + &buffers, |
| 278 | + &serialized.children, |
| 279 | + ), |
| 280 | + &session, |
| 281 | + )?; |
| 282 | + assert_arrays_eq!(decoded, values, &mut ctx); |
| 283 | + Ok(()) |
| 284 | + } |
| 285 | +} |
0 commit comments