vortex_array/arrays/decimal/vtable/
mod.rs1use std::hash::Hasher;
5
6use prost::Message;
7use vortex_buffer::Alignment;
8use vortex_error::VortexResult;
9use vortex_error::vortex_bail;
10use vortex_error::vortex_ensure;
11use vortex_error::vortex_panic;
12use vortex_session::VortexSession;
13
14use crate::ArrayParts;
15use crate::ArrayRef;
16use crate::ExecutionCtx;
17use crate::ExecutionResult;
18use crate::array::Array;
19use crate::array::ArrayView;
20use crate::array::VTable;
21use crate::arrays::decimal::DecimalData;
22use crate::buffer::BufferHandle;
23use crate::builders::ArrayBuilder;
24use crate::builders::DecimalBuilder;
25use crate::dtype::DType;
26use crate::dtype::DecimalType;
27use crate::dtype::NativeDecimalType;
28use crate::match_each_decimal_value_type;
29use crate::serde::ArrayChildren;
30use crate::validity::Validity;
31mod kernel;
32mod operations;
33mod validity;
34
35use std::hash::Hash;
36
37use vortex_session::registry::CachedId;
38
39use crate::EqMode;
40use crate::array::ArrayId;
41use crate::arrays::decimal::array::SLOT_NAMES;
42use crate::arrays::decimal::compute::rules::RULES;
43use crate::hash::ArrayEq;
44use crate::hash::ArrayHash;
45pub type DecimalArray = Array<Decimal>;
47
48pub(crate) fn initialize(session: &VortexSession) {
49 kernel::initialize(session);
50}
51
52#[derive(prost::Message)]
54pub struct DecimalMetadata {
55 #[prost(enumeration = "DecimalType", tag = "1")]
56 pub(super) values_type: i32,
57}
58
59impl ArrayHash for DecimalData {
60 fn array_hash<H: Hasher>(&self, state: &mut H, accuracy: EqMode) {
61 self.values.array_hash(state, accuracy);
62 std::mem::discriminant(&self.values_type).hash(state);
63 }
64}
65
66impl ArrayEq for DecimalData {
67 fn array_eq(&self, other: &Self, accuracy: EqMode) -> bool {
68 self.values.array_eq(&other.values, accuracy) && self.values_type == other.values_type
69 }
70}
71
72impl VTable for Decimal {
73 type TypedArrayData = DecimalData;
74
75 type OperationsVTable = Self;
76 type ValidityVTable = Self;
77
78 fn id(&self) -> ArrayId {
79 static ID: CachedId = CachedId::new("vortex.decimal");
80 *ID
81 }
82
83 fn nbuffers(_array: ArrayView<'_, Self>) -> usize {
84 1
85 }
86
87 fn buffer(array: ArrayView<'_, Self>, idx: usize) -> BufferHandle {
88 match idx {
89 0 => array.values.clone(),
90 _ => vortex_panic!("DecimalArray buffer index {idx} out of bounds"),
91 }
92 }
93
94 fn buffer_name(_array: ArrayView<'_, Self>, idx: usize) -> Option<String> {
95 match idx {
96 0 => Some("values".to_string()),
97 _ => None,
98 }
99 }
100
101 fn with_buffers(
102 &self,
103 array: ArrayView<'_, Self>,
104 buffers: &[BufferHandle],
105 ) -> VortexResult<ArrayParts<Self>> {
106 vortex_ensure!(
107 buffers.len() == 1,
108 "Expected 1 buffer, got {}",
109 buffers.len()
110 );
111 let mut data = array.data().clone();
112 data.values = buffers[0].clone();
113 Ok(
114 ArrayParts::new(self.clone(), array.dtype().clone(), array.len(), data)
115 .with_slots(array.slots().iter().cloned().collect()),
116 )
117 }
118
119 fn serialize(
120 array: ArrayView<'_, Self>,
121 _session: &VortexSession,
122 ) -> VortexResult<Option<Vec<u8>>> {
123 Ok(Some(
124 DecimalMetadata {
125 values_type: array.values_type() as i32,
126 }
127 .encode_to_vec(),
128 ))
129 }
130
131 fn validate(
132 &self,
133 data: &DecimalData,
134 dtype: &DType,
135 len: usize,
136 slots: &[Option<ArrayRef>],
137 ) -> VortexResult<()> {
138 let DType::Decimal(_, nullability) = dtype else {
139 vortex_bail!("Expected decimal dtype, got {dtype:?}");
140 };
141 vortex_ensure!(
142 data.len() == len,
143 InvalidArgument:
144 "DecimalArray length {} does not match outer length {}",
145 data.len(),
146 len
147 );
148 let validity = crate::array::child_to_validity(slots[0].as_ref(), *nullability);
149 if let Some(validity_len) = validity.maybe_len() {
150 vortex_ensure!(
151 validity_len == len,
152 InvalidArgument:
153 "DecimalArray validity len {} does not match outer length {}",
154 validity_len,
155 len
156 );
157 }
158
159 Ok(())
160 }
161
162 fn deserialize(
163 &self,
164 dtype: &DType,
165 len: usize,
166 metadata: &[u8],
167 buffers: &[BufferHandle],
168 children: &dyn ArrayChildren,
169 _session: &VortexSession,
170 ) -> VortexResult<ArrayParts<Self>> {
171 let metadata = DecimalMetadata::decode(metadata)?;
172 if buffers.len() != 1 {
173 vortex_bail!("Expected 1 buffer, got {}", buffers.len());
174 }
175 let values = buffers[0].clone();
176
177 let validity = if children.is_empty() {
178 Validity::from(dtype.nullability())
179 } else if children.len() == 1 {
180 let validity = children.get(0, &Validity::DTYPE, len)?;
181 Validity::Array(validity)
182 } else {
183 vortex_bail!("Expected 0 or 1 child, got {}", children.len());
184 };
185
186 let Some(decimal_dtype) = dtype.as_decimal_opt() else {
187 vortex_bail!("Expected Decimal dtype, got {:?}", dtype)
188 };
189
190 let slots = DecimalData::make_slots(&validity, len);
191 let data = match_each_decimal_value_type!(metadata.values_type(), |D| {
192 vortex_ensure!(
194 values.is_aligned_to(Alignment::of::<D>()),
195 "DecimalArray buffer not aligned for values type {:?}",
196 D::DECIMAL_TYPE
197 );
198 DecimalData::try_new_handle(values, metadata.values_type(), *decimal_dtype)
199 })?;
200 Ok(ArrayParts::new(self.clone(), dtype.clone(), len, data).with_slots(slots))
201 }
202
203 fn slot_name(_array: ArrayView<'_, Self>, idx: usize) -> String {
204 SLOT_NAMES[idx].to_string()
205 }
206
207 fn execute(array: Array<Self>, _ctx: &mut ExecutionCtx) -> VortexResult<ExecutionResult> {
208 Ok(ExecutionResult::done(array))
209 }
210
211 fn append_to_builder(
212 array: ArrayView<'_, Self>,
213 builder: &mut dyn ArrayBuilder,
214 ctx: &mut ExecutionCtx,
215 ) -> VortexResult<()> {
216 let Some(builder) = builder.as_any_mut().downcast_mut::<DecimalBuilder>() else {
217 vortex_bail!("append_to_builder for Decimal requires a DecimalBuilder");
218 };
219 builder.append_decimal_array(&array.into_owned(), ctx)
220 }
221
222 fn reduce_parent(
223 array: ArrayView<'_, Self>,
224 parent: &ArrayRef,
225 child_idx: usize,
226 ) -> VortexResult<Option<ArrayRef>> {
227 RULES.evaluate(array, parent, child_idx)
228 }
229}
230
231#[derive(Clone, Debug)]
232pub struct Decimal;
233
234#[cfg(test)]
235mod tests {
236 use vortex_buffer::ByteBufferMut;
237 use vortex_buffer::buffer;
238 use vortex_session::registry::ReadContext;
239
240 use crate::ArrayContext;
241 use crate::IntoArray;
242 use crate::VortexSessionExecute;
243 use crate::array_session;
244 use crate::arrays::Decimal;
245 use crate::arrays::DecimalArray;
246 use crate::assert_arrays_eq;
247 use crate::dtype::DecimalDType;
248 use crate::serde::SerializeOptions;
249 use crate::serde::SerializedArray;
250 use crate::validity::Validity;
251
252 #[test]
253 fn test_array_serde() {
254 let session = array_session();
255 let array = DecimalArray::new(
256 buffer![100i128, 200i128, 300i128, 400i128, 500i128],
257 DecimalDType::new(10, 2),
258 Validity::NonNullable,
259 );
260 let dtype = array.dtype().clone();
261
262 let array_ctx = ArrayContext::empty();
263 let out = array
264 .into_array()
265 .serialize(&array_ctx, &session, &SerializeOptions::default())
266 .unwrap();
267 let mut concat = ByteBufferMut::empty();
269 for buf in out {
270 concat.extend_from_slice(buf.as_ref());
271 }
272
273 let concat = concat.freeze();
274
275 let parts = SerializedArray::try_from(concat).unwrap();
276 let decoded = parts
277 .decode(&dtype, 5, &ReadContext::new(array_ctx.to_ids()), &session)
278 .unwrap();
279 assert!(decoded.is::<Decimal>());
280 }
281
282 #[test]
283 fn test_nullable_decimal_serde_roundtrip() {
284 let session = array_session();
285 let mut ctx = session.create_execution_ctx();
286 let array = DecimalArray::new(
287 buffer![1234567i32, 0i32, -9999999i32],
288 DecimalDType::new(7, 3),
289 Validity::from_iter([true, false, true]),
290 );
291 let dtype = array.dtype().clone();
292 let len = array.len();
293
294 let array_ctx = ArrayContext::empty();
295 let out = array
296 .clone()
297 .into_array()
298 .serialize(&array_ctx, &session, &SerializeOptions::default())
299 .unwrap();
300 let mut concat = ByteBufferMut::empty();
301 for buf in out {
302 concat.extend_from_slice(buf.as_ref());
303 }
304
305 let parts = SerializedArray::try_from(concat.freeze()).unwrap();
306 let decoded = parts
307 .decode(&dtype, len, &ReadContext::new(array_ctx.to_ids()), &session)
308 .unwrap();
309
310 assert_arrays_eq!(decoded, array, &mut ctx);
311 }
312}