1use std::collections::BTreeMap;
4use std::sync::Arc;
5
6use arrow::compute::take;
7use arrow_array::{ArrayRef, RecordBatch, UInt32Array};
8use arrow_row::{RowConverter, SortField};
9
10use crate::state::{key_hash, NodeId, VnodeAssignmentSnapshot};
11
12pub const ROUTE_TARGET_BATCH_BYTES: usize = 4 * 1024 * 1024;
15pub const ROUTE_MAX_BATCH_BYTES: usize = 8 * 1024 * 1024;
17pub const ROUTE_MAX_BATCH_ROWS: usize = 65_536;
19
20pub(crate) fn logical_batch_bytes(batch: &RecordBatch) -> Result<usize, arrow_schema::ArrowError> {
24 batch.columns().iter().try_fold(0usize, |total, column| {
25 let data = column.to_data();
26 let bytes = data
27 .get_slice_memory_size()?
28 .checked_add(variadic_buffer_bytes(&data)?)
29 .ok_or_else(|| {
30 arrow_schema::ArrowError::ComputeError(
31 "logical shuffle batch size overflow".to_string(),
32 )
33 })?;
34 total.checked_add(bytes).ok_or_else(|| {
35 arrow_schema::ArrowError::ComputeError(
36 "logical shuffle batch size overflow".to_string(),
37 )
38 })
39 })
40}
41
42fn variadic_buffer_bytes(
43 data: &arrow::array::ArrayData,
44) -> Result<usize, arrow_schema::ArrowError> {
45 let own = if matches!(
46 data.data_type(),
47 arrow_schema::DataType::Utf8View | arrow_schema::DataType::BinaryView
48 ) {
49 data.buffers()
50 .iter()
51 .skip(1)
52 .try_fold(0usize, |total, buffer| {
53 total.checked_add(buffer.len()).ok_or_else(|| {
54 arrow_schema::ArrowError::ComputeError(
55 "variadic Arrow buffer size overflow".to_string(),
56 )
57 })
58 })?
59 } else {
60 0
61 };
62 data.child_data().iter().try_fold(own, |total, child| {
63 total
64 .checked_add(variadic_buffer_bytes(child)?)
65 .ok_or_else(|| {
66 arrow_schema::ArrowError::ComputeError(
67 "nested variadic Arrow buffer size overflow".to_string(),
68 )
69 })
70 })
71}
72
73#[derive(Debug, Clone)]
75pub struct LocalRoute {
76 pub vnode: u32,
78 pub batch: RecordBatch,
80}
81
82#[derive(Debug, Clone)]
84pub struct RemoteRoute {
85 pub owner: NodeId,
87 pub routed_vnodes: Arc<[u32]>,
89 pub batch: RecordBatch,
91}
92
93#[derive(Debug, Clone, Default)]
95pub struct CheckpointRoutePlan {
96 pub local: Vec<LocalRoute>,
98 pub remote: Vec<RemoteRoute>,
100}
101
102#[derive(Debug, thiserror::Error)]
104pub enum ShuffleRoutingError {
105 #[error("shuffle routing requires at least one key column")]
107 EmptyKey,
108 #[error("shuffle routing requires a nonzero vnode count")]
110 EmptyVnodeSpace,
111 #[error("shuffle key column {index} is outside the {columns}-column batch")]
113 KeyColumnOutOfRange {
114 index: usize,
116 columns: usize,
118 },
119 #[error("shuffle key column {index} has unsupported partition type {data_type}")]
121 UnsupportedKeyType {
122 index: usize,
124 data_type: arrow_schema::DataType,
126 },
127 #[error("shuffle Arrow routing: {0}")]
129 Arrow(#[from] arrow_schema::ArrowError),
130 #[error("shuffle route has {vnodes} vnode entries for {rows} rows")]
132 RowVnodeCardinality {
133 rows: usize,
135 vnodes: usize,
137 },
138 #[error("shuffle batch has {rows} rows; maximum supported is {}", u32::MAX)]
140 TooManyRows {
141 rows: usize,
143 },
144 #[error("shuffle vnode {vnode} is outside the {vnode_count}-vnode assignment")]
146 VnodeOutOfRange {
147 vnode: u32,
149 vnode_count: usize,
151 },
152 #[error("shuffle vnode {vnode} is unassigned")]
154 UnassignedVnode {
155 vnode: u32,
157 },
158 #[error("shuffle row for vnode {vnode} occupies {bytes} decoded bytes; hard limit is {limit}")]
160 OversizedRow {
161 vnode: u32,
163 bytes: usize,
165 limit: usize,
167 },
168}
169
170impl ShuffleRoutingError {
171 #[must_use]
173 pub const fn is_not_ready(&self) -> bool {
174 matches!(self, Self::UnassignedVnode { .. })
175 }
176}
177
178pub fn row_vnodes(
183 batch: &RecordBatch,
184 columns: &[usize],
185 vnode_count: u32,
186) -> Result<Vec<u32>, ShuffleRoutingError> {
187 if columns.is_empty() {
188 return Err(ShuffleRoutingError::EmptyKey);
189 }
190 if vnode_count == 0 {
191 return Err(ShuffleRoutingError::EmptyVnodeSpace);
192 }
193 let cols: Vec<ArrayRef> = columns
194 .iter()
195 .map(|&index| {
196 let column = batch.columns().get(index).cloned().ok_or(
197 ShuffleRoutingError::KeyColumnOutOfRange {
198 index,
199 columns: batch.num_columns(),
200 },
201 )?;
202 if !is_supported_key_type(column.data_type()) {
203 return Err(ShuffleRoutingError::UnsupportedKeyType {
204 index,
205 data_type: column.data_type().clone(),
206 });
207 }
208 Ok(column)
209 })
210 .collect::<Result<_, _>>()?;
211 let fields: Vec<SortField> = cols
212 .iter()
213 .map(|column| SortField::new(column.data_type().clone()))
214 .collect();
215 let converter = RowConverter::new(fields)?;
216 let rows = converter.convert_columns(&cols)?;
217 (0..batch.num_rows())
218 .map(|row| {
219 u32::try_from(key_hash(rows.row(row).as_ref()) % u64::from(vnode_count))
220 .map_err(|_| ShuffleRoutingError::EmptyVnodeSpace)
221 })
222 .collect()
223}
224
225fn is_supported_key_type(data_type: &arrow_schema::DataType) -> bool {
226 match data_type {
227 arrow_schema::DataType::Dictionary(indices, values) => {
230 matches!(
231 indices.as_ref(),
232 arrow_schema::DataType::Int8
233 | arrow_schema::DataType::Int16
234 | arrow_schema::DataType::Int32
235 | arrow_schema::DataType::Int64
236 | arrow_schema::DataType::UInt8
237 | arrow_schema::DataType::UInt16
238 | arrow_schema::DataType::UInt32
239 | arrow_schema::DataType::UInt64
240 ) && is_supported_key_type(values)
241 }
242 arrow_schema::DataType::RunEndEncoded(_, _) => false,
245 data_type => !data_type.is_floating() && !data_type.is_nested(),
246 }
247}
248
249pub fn route_checkpointed_batch(
259 batch: &RecordBatch,
260 row_vnodes: &[u32],
261 assignment: &VnodeAssignmentSnapshot,
262 self_id: NodeId,
263) -> Result<CheckpointRoutePlan, ShuffleRoutingError> {
264 if row_vnodes.len() != batch.num_rows() {
265 return Err(ShuffleRoutingError::RowVnodeCardinality {
266 rows: batch.num_rows(),
267 vnodes: row_vnodes.len(),
268 });
269 }
270 if batch.num_rows() > usize::try_from(u32::MAX).unwrap_or(usize::MAX) {
271 return Err(ShuffleRoutingError::TooManyRows {
272 rows: batch.num_rows(),
273 });
274 }
275 if batch.num_rows() == 0 {
276 return Ok(CheckpointRoutePlan::default());
277 }
278
279 let mut local_groups: BTreeMap<u32, Vec<u32>> = BTreeMap::new();
280 let mut remote_groups: BTreeMap<NodeId, Vec<(u32, u32)>> = BTreeMap::new();
281 for (row, &vnode) in row_vnodes.iter().enumerate() {
282 let owner = assignment
283 .owners()
284 .get(usize::try_from(vnode).unwrap_or(usize::MAX))
285 .copied()
286 .ok_or(ShuffleRoutingError::VnodeOutOfRange {
287 vnode,
288 vnode_count: assignment.owners().len(),
289 })?;
290 if owner.is_unassigned() {
291 return Err(ShuffleRoutingError::UnassignedVnode { vnode });
292 }
293 let row = u32::try_from(row).map_err(|_| ShuffleRoutingError::TooManyRows {
294 rows: batch.num_rows(),
295 })?;
296 if owner == self_id {
297 local_groups.entry(vnode).or_default().push(row);
298 } else {
299 remote_groups.entry(owner).or_default().push((row, vnode));
300 }
301 }
302
303 let mut plan = CheckpointRoutePlan::default();
304 for (vnode, indices) in local_groups {
305 for (_, slice) in bounded_slices(batch, &indices, vnode)? {
306 plan.local.push(LocalRoute {
307 vnode,
308 batch: slice,
309 });
310 }
311 }
312 for (owner, rows) in remote_groups {
313 let indices: Vec<u32> = rows.iter().map(|(row, _)| *row).collect();
314 for (range, slice) in bounded_slices(batch, &indices, rows[0].1)? {
315 let mut routed_vnodes: Vec<u32> = rows[range].iter().map(|(_, vnode)| *vnode).collect();
316 routed_vnodes.sort_unstable();
317 routed_vnodes.dedup();
318 plan.remote.push(RemoteRoute {
319 owner,
320 routed_vnodes: routed_vnodes.into(),
321 batch: slice,
322 });
323 }
324 }
325 debug_assert_eq!(
326 plan.local
327 .iter()
328 .map(|route| route.batch.num_rows())
329 .sum::<usize>()
330 + plan
331 .remote
332 .iter()
333 .map(|route| route.batch.num_rows())
334 .sum::<usize>(),
335 batch.num_rows()
336 );
337 Ok(plan)
338}
339
340fn bounded_slices(
341 batch: &RecordBatch,
342 indices: &[u32],
343 vnode: u32,
344) -> Result<Vec<(std::ops::Range<usize>, RecordBatch)>, ShuffleRoutingError> {
345 let estimated_row_bytes = logical_batch_bytes(batch)?
346 .div_ceil(batch.num_rows())
347 .max(1);
348 let initial_rows =
349 ROUTE_MAX_BATCH_ROWS.min((ROUTE_TARGET_BATCH_BYTES / estimated_row_bytes).max(1));
350 let mut chunks = Vec::new();
351 let mut start = 0;
352 while start < indices.len() {
353 let mut end = (start + initial_rows).min(indices.len());
354 loop {
355 let slice = take_rows(batch, &indices[start..end])?;
356 let bytes = logical_batch_bytes(&slice)?;
357 let rows = end - start;
358 if bytes <= ROUTE_TARGET_BATCH_BYTES || rows == 1 {
359 if bytes > ROUTE_MAX_BATCH_BYTES {
360 return Err(ShuffleRoutingError::OversizedRow {
361 vnode,
362 bytes,
363 limit: ROUTE_MAX_BATCH_BYTES,
364 });
365 }
366 chunks.push((start..end, slice));
367 start = end;
368 break;
369 }
370 end = start + (rows / 2).max(1);
371 }
372 }
373 Ok(chunks)
374}
375
376fn take_rows(batch: &RecordBatch, indices: &[u32]) -> Result<RecordBatch, ShuffleRoutingError> {
377 let indices = UInt32Array::from(indices.to_vec());
378 let columns = batch
379 .columns()
380 .iter()
381 .map(|column| take(column, &indices, None))
382 .collect::<Result<Vec<_>, _>>()?;
383 Ok(RecordBatch::try_new(batch.schema(), columns)?)
384}
385
386#[cfg(test)]
387mod tests {
388 use arrow_array::types::{Int16Type, Int8Type};
389 use arrow_array::{
390 Array as _, DictionaryArray, Float64Array, Int16Array, Int64Array, Int8Array, RecordBatch,
391 StringArray, StringViewArray,
392 };
393 use arrow_schema::{DataType, Field, Schema, UnionFields, UnionMode};
394
395 use super::*;
396 use crate::state::VnodeRegistry;
397
398 fn values(values: &[i64]) -> RecordBatch {
399 RecordBatch::try_new(
400 Arc::new(Schema::new(vec![Field::new(
401 "value",
402 DataType::Int64,
403 false,
404 )])),
405 vec![Arc::new(Int64Array::from(values.to_vec()))],
406 )
407 .unwrap()
408 }
409
410 #[test]
411 fn route_plan_uses_pinned_assignment_and_conserves_rows() {
412 let registry = VnodeRegistry::single_owner(2, NodeId(1));
413 let pinned = registry.versioned_snapshot();
414 registry.set_assignment(Arc::from([NodeId(2), NodeId(2)]));
415
416 let plan =
417 route_checkpointed_batch(&values(&[10, 20]), &[0, 1], &pinned, NodeId(1)).unwrap();
418
419 assert_eq!(pinned.version(), 1);
420 assert_eq!(registry.assignment_version(), 2);
421 assert_eq!(plan.local.len(), 2);
422 assert!(plan.remote.is_empty());
423 assert_eq!(
424 plan.local
425 .iter()
426 .map(|route| route.batch.num_rows())
427 .sum::<usize>(),
428 2
429 );
430 }
431
432 #[test]
433 fn unassigned_vnode_fails_before_any_plan_is_returned() {
434 let registry = VnodeRegistry::single_owner(2, NodeId(1));
435 registry.set_assignment(Arc::from([NodeId(1), NodeId::UNASSIGNED]));
436 let assignment = registry.versioned_snapshot();
437
438 let error = route_checkpointed_batch(&values(&[10, 20]), &[0, 1], &assignment, NodeId(1))
439 .unwrap_err();
440
441 assert!(error.is_not_ready());
442 }
443
444 #[test]
445 fn remote_routes_preserve_user_vnode_named_field() {
446 let schema = Arc::new(Schema::new(vec![
447 Field::new("__laminar_vnode", DataType::Utf8, false),
448 Field::new("value", DataType::Int64, false),
449 ]));
450 let batch = RecordBatch::try_new(
451 schema.clone(),
452 vec![
453 Arc::new(StringArray::from(vec!["user-a", "user-b"])),
454 Arc::new(Int64Array::from(vec![10, 20])),
455 ],
456 )
457 .unwrap();
458 let registry = VnodeRegistry::single_owner(2, NodeId(2));
459
460 let plan =
461 route_checkpointed_batch(&batch, &[0, 1], ®istry.versioned_snapshot(), NodeId(1))
462 .unwrap();
463
464 assert_eq!(plan.remote.len(), 1);
465 assert_eq!(plan.remote[0].routed_vnodes.as_ref(), &[0, 1]);
466 assert_eq!(plan.remote[0].batch.schema(), schema);
467 assert_eq!(plan.remote[0].batch.num_columns(), 2);
468 }
469
470 #[test]
471 fn large_owner_group_is_split_below_the_decoded_bound() {
472 let payload = "x".repeat(1024);
473 let rows = ROUTE_TARGET_BATCH_BYTES / payload.len() + 512;
474 let batch = RecordBatch::try_new(
475 Arc::new(Schema::new(vec![Field::new(
476 "value",
477 DataType::Utf8,
478 false,
479 )])),
480 vec![Arc::new(StringArray::from(vec![payload; rows]))],
481 )
482 .unwrap();
483 let registry = VnodeRegistry::single_owner(1, NodeId(2));
484 let plan = route_checkpointed_batch(
485 &batch,
486 &vec![0; rows],
487 ®istry.versioned_snapshot(),
488 NodeId(1),
489 )
490 .unwrap();
491
492 assert!(plan.remote.len() > 1);
493 assert_eq!(
494 plan.remote
495 .iter()
496 .map(|route| route.batch.num_rows())
497 .sum::<usize>(),
498 rows
499 );
500 assert!(plan.remote.iter().all(|route| {
501 route.batch.num_rows() <= ROUTE_MAX_BATCH_ROWS
502 && logical_batch_bytes(&route.batch).unwrap() <= ROUTE_MAX_BATCH_BYTES
503 }));
504 }
505
506 #[test]
507 fn string_view_variadic_buffers_count_toward_the_hard_bound() {
508 let value = "x".repeat(129);
509 let array = StringViewArray::from_iter_values(std::iter::repeat_n(
510 value.as_str(),
511 ROUTE_MAX_BATCH_ROWS,
512 ));
513 let batch = RecordBatch::try_new(
514 Arc::new(Schema::new(vec![Field::new(
515 "value",
516 DataType::Utf8View,
517 false,
518 )])),
519 vec![Arc::new(array)],
520 )
521 .unwrap();
522
523 assert!(logical_batch_bytes(&batch).unwrap() > ROUTE_MAX_BATCH_BYTES);
524 }
525
526 #[test]
527 fn narrow_slices_of_shared_backing_use_referenced_bytes() {
528 let backing = Int64Array::from(vec![1; 2_000_000]);
529 let column_count = 256;
530 let fields = (0..column_count)
531 .map(|index| Field::new(format!("c{index}"), DataType::Int64, false))
532 .collect::<Vec<_>>();
533 let columns = (0..column_count)
534 .map(|index| Arc::new(backing.slice(index, 1)) as arrow_array::ArrayRef)
535 .collect::<Vec<_>>();
536 let batch = RecordBatch::try_new(Arc::new(Schema::new(fields)), columns).unwrap();
537
538 assert!(batch.get_array_memory_size() > ROUTE_MAX_BATCH_BYTES);
539 assert_eq!(
540 logical_batch_bytes(&batch).unwrap(),
541 column_count * std::mem::size_of::<i64>()
542 );
543 }
544
545 #[test]
546 fn mixed_routes_have_deterministic_owner_order_and_conserve_input_order() {
547 let registry = VnodeRegistry::single_owner(4, NodeId(1));
548 registry.set_assignment(Arc::from([NodeId(2), NodeId(1), NodeId(2), NodeId(3)]));
549 let batch = values(&[10, 20, 30, 40, 50, 60]);
550
551 let plan = route_checkpointed_batch(
552 &batch,
553 &[2, 0, 3, 2, 1, 0],
554 ®istry.versioned_snapshot(),
555 NodeId(1),
556 )
557 .unwrap();
558
559 assert_eq!(plan.local.len(), 1);
560 assert_eq!(plan.local[0].vnode, 1);
561 assert_eq!(
562 plan.local[0]
563 .batch
564 .column(0)
565 .as_any()
566 .downcast_ref::<Int64Array>()
567 .unwrap()
568 .values(),
569 &[50]
570 );
571 assert_eq!(
572 plan.remote
573 .iter()
574 .map(|route| route.owner)
575 .collect::<Vec<_>>(),
576 vec![NodeId(2), NodeId(3)]
577 );
578 assert_eq!(plan.remote[0].routed_vnodes.as_ref(), &[0, 2]);
579 assert_eq!(
580 plan.remote[0]
581 .batch
582 .column(0)
583 .as_any()
584 .downcast_ref::<Int64Array>()
585 .unwrap()
586 .values(),
587 &[10, 20, 40, 60]
588 );
589 }
590
591 #[test]
592 fn single_oversized_row_is_terminal_before_a_plan_is_returned() {
593 let batch = RecordBatch::try_new(
594 Arc::new(Schema::new(vec![Field::new(
595 "value",
596 DataType::Utf8,
597 false,
598 )])),
599 vec![Arc::new(StringArray::from(vec![
600 "x".repeat(ROUTE_MAX_BATCH_BYTES + 1)
601 ]))],
602 )
603 .unwrap();
604 let registry = VnodeRegistry::single_owner(1, NodeId(2));
605
606 let error =
607 route_checkpointed_batch(&batch, &[0], ®istry.versioned_snapshot(), NodeId(1))
608 .unwrap_err();
609
610 assert!(matches!(error, ShuffleRoutingError::OversizedRow { .. }));
611 assert!(!error.is_not_ready());
612 }
613
614 #[test]
615 fn row_hashing_rejects_invalid_dimensions_without_panicking() {
616 let batch = values(&[1]);
617 assert!(matches!(
618 row_vnodes(&batch, &[], 1),
619 Err(ShuffleRoutingError::EmptyKey)
620 ));
621 assert!(matches!(
622 row_vnodes(&batch, &[0], 0),
623 Err(ShuffleRoutingError::EmptyVnodeSpace)
624 ));
625 assert!(matches!(
626 row_vnodes(&batch, &[1], 1),
627 Err(ShuffleRoutingError::KeyColumnOutOfRange { .. })
628 ));
629 }
630
631 #[test]
632 fn partitioning_abi_v1_arrow_row_golden_vectors() {
633 const KEY_GROUPS: u32 = 257;
634
635 let strings = RecordBatch::try_new(
636 Arc::new(Schema::new(vec![Field::new("key", DataType::Utf8, true)])),
637 vec![Arc::new(StringArray::from(vec![
638 None,
639 Some(""),
640 Some("alpha"),
641 Some("snowman-☃"),
642 ]))],
643 )
644 .unwrap();
645 let string_groups = row_vnodes(&strings, &[0], KEY_GROUPS).unwrap();
646
647 let integers = RecordBatch::try_new(
648 Arc::new(Schema::new(vec![Field::new("key", DataType::Int64, true)])),
649 vec![Arc::new(Int64Array::from(vec![
650 Some(i64::MIN),
651 Some(-1),
652 None,
653 Some(0),
654 Some(1),
655 Some(i64::MAX),
656 ]))],
657 )
658 .unwrap();
659 let integer_groups = row_vnodes(&integers, &[0], KEY_GROUPS).unwrap();
660
661 let composite = RecordBatch::try_new(
662 Arc::new(Schema::new(vec![
663 Field::new("tenant", DataType::Utf8, true),
664 Field::new("account", DataType::Int64, true),
665 ])),
666 vec![
667 Arc::new(StringArray::from(vec![Some("a"), Some("a"), None, None])),
668 Arc::new(Int64Array::from(vec![Some(1), None, Some(1), None])),
669 ],
670 )
671 .unwrap();
672 let composite_groups = row_vnodes(&composite, &[0, 1], KEY_GROUPS).unwrap();
673
674 let dictionary =
675 DictionaryArray::<Int8Type>::from_iter([Some("alpha"), None, Some(""), Some("alpha")]);
676 let dictionary = RecordBatch::try_new(
677 Arc::new(Schema::new(vec![Field::new(
678 "key",
679 dictionary.data_type().clone(),
680 true,
681 )])),
682 vec![Arc::new(dictionary)],
683 )
684 .unwrap();
685 let dictionary_groups = row_vnodes(&dictionary, &[0], KEY_GROUPS).unwrap();
686
687 assert_eq!(
688 (
689 string_groups,
690 integer_groups,
691 composite_groups,
692 dictionary_groups,
693 ),
694 (
695 vec![211, 224, 44, 94],
696 vec![111, 202, 114, 90, 180, 32],
697 vec![26, 208, 118, 52],
698 vec![44, 211, 224, 44],
699 )
700 );
701 }
702
703 #[test]
704 fn partitioning_abi_v1_rejects_non_scalar_and_floating_point_keys() {
705 let batch = RecordBatch::try_new(
706 Arc::new(Schema::new(vec![Field::new(
707 "key",
708 DataType::Float64,
709 false,
710 )])),
711 vec![Arc::new(Float64Array::from(vec![0.0, -0.0, f64::NAN]))],
712 )
713 .unwrap();
714
715 assert!(matches!(
716 row_vnodes(&batch, &[0], 257),
717 Err(ShuffleRoutingError::UnsupportedKeyType {
718 index: 0,
719 data_type: DataType::Float64,
720 })
721 ));
722
723 let item = Arc::new(Field::new("item", DataType::Int64, true));
724 let nested = [
725 DataType::List(Arc::clone(&item)),
726 DataType::ListView(Arc::clone(&item)),
727 DataType::LargeList(Arc::clone(&item)),
728 DataType::LargeListView(Arc::clone(&item)),
729 DataType::FixedSizeList(Arc::clone(&item), 2),
730 DataType::Struct(vec![Field::new("item", DataType::Int64, true)].into()),
731 DataType::Union(
732 UnionFields::try_new([0], [Field::new("item", DataType::Int64, true)]).unwrap(),
733 UnionMode::Sparse,
734 ),
735 DataType::Map(
736 Arc::new(Field::new(
737 "entries",
738 DataType::Struct(
739 vec![
740 Field::new("key", DataType::Utf8, false),
741 Field::new("value", DataType::Int64, true),
742 ]
743 .into(),
744 ),
745 false,
746 )),
747 false,
748 ),
749 ];
750 assert!(nested
751 .iter()
752 .all(|data_type| !is_supported_key_type(data_type)));
753 assert!(!is_supported_key_type(&DataType::Dictionary(
754 Box::new(DataType::Int8),
755 Box::new(DataType::Float64),
756 )));
757 assert!(!is_supported_key_type(&DataType::Dictionary(
758 Box::new(DataType::Int8),
759 Box::new(DataType::List(item)),
760 )));
761 assert!(!is_supported_key_type(&DataType::Dictionary(
762 Box::new(DataType::Float64),
763 Box::new(DataType::Utf8),
764 )));
765 }
766
767 #[test]
768 fn dictionary_encoding_does_not_change_partitioning() {
769 const KEY_GROUPS: u32 = 257;
770
771 let int8 = DictionaryArray::<Int8Type>::try_new(
772 Int8Array::from(vec![Some(0), None, Some(1), Some(0)]),
773 Arc::new(StringArray::from(vec!["alpha", ""])),
774 )
775 .unwrap();
776 let int16 = DictionaryArray::<Int16Type>::try_new(
777 Int16Array::from(vec![Some(2), None, Some(0), Some(2)]),
778 Arc::new(StringArray::from(vec!["", "unused", "alpha"])),
779 )
780 .unwrap();
781
782 let int8 = RecordBatch::try_new(
783 Arc::new(Schema::new(vec![Field::new(
784 "key",
785 int8.data_type().clone(),
786 true,
787 )])),
788 vec![Arc::new(int8)],
789 )
790 .unwrap();
791 let int16 = RecordBatch::try_new(
792 Arc::new(Schema::new(vec![Field::new(
793 "key",
794 int16.data_type().clone(),
795 true,
796 )])),
797 vec![Arc::new(int16)],
798 )
799 .unwrap();
800
801 let int8_groups = row_vnodes(&int8, &[0], KEY_GROUPS).unwrap();
802 let int16_groups = row_vnodes(&int16, &[0], KEY_GROUPS).unwrap();
803 assert_eq!(int8_groups, vec![44, 211, 224, 44]);
804 assert_eq!(int16_groups, int8_groups);
805 }
806}