Skip to main content

laminar_core/shuffle/
routing.rs

1//! Row-to-vnode routing shared by checkpointed cluster shuffle paths.
2
3use 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
12/// Target decoded size for one routed batch. This is intentionally below the hard receiver bound
13/// so schema/allocator variance does not turn ordinary skew into a transport failure.
14pub const ROUTE_TARGET_BATCH_BYTES: usize = 4 * 1024 * 1024;
15/// Hard decoded size admitted for one logical shuffle batch.
16pub const ROUTE_MAX_BATCH_BYTES: usize = 8 * 1024 * 1024;
17/// Row-count bound prevents tiny or zero-width rows from producing unbounded logical frames.
18pub const ROUTE_MAX_BATCH_ROWS: usize = 65_536;
19
20/// Logical Arrow bytes referenced by this slice, independent of backing-buffer capacity or owner.
21/// This is the stable bound for IPC content; transport reservations account the retained backing
22/// allocation separately.
23pub(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/// A local vnode slice in deterministic vnode order.
74#[derive(Debug, Clone)]
75pub struct LocalRoute {
76    /// Exact local vnode.
77    pub vnode: u32,
78    /// Rows for that vnode, preserving their input order.
79    pub batch: RecordBatch,
80}
81
82/// One bounded owner-coalesced remote batch.
83#[derive(Debug, Clone)]
84pub struct RemoteRoute {
85    /// Certified remote owner.
86    pub owner: NodeId,
87    /// Ascending, duplicate-free vnode set represented by `batch`.
88    pub routed_vnodes: Arc<[u32]>,
89    /// Rows for this owner, preserving their input order.
90    pub batch: RecordBatch,
91}
92
93/// Complete, conserving routing plan staged before any remote send.
94#[derive(Debug, Clone, Default)]
95pub struct CheckpointRoutePlan {
96    /// Local chunks, ordered by vnode and then input row.
97    pub local: Vec<LocalRoute>,
98    /// Remote chunks, ordered by owner and then input row.
99    pub remote: Vec<RemoteRoute>,
100}
101
102/// A routing failure. No plan is returned, so callers cannot partially send or silently drop rows.
103#[derive(Debug, thiserror::Error)]
104pub enum ShuffleRoutingError {
105    /// Keyed routing requires at least one key column.
106    #[error("shuffle routing requires at least one key column")]
107    EmptyKey,
108    /// Modulo routing is undefined for an empty vnode space.
109    #[error("shuffle routing requires a nonzero vnode count")]
110    EmptyVnodeSpace,
111    /// A resolved key index no longer exists in the batch schema.
112    #[error("shuffle key column {index} is outside the {columns}-column batch")]
113    KeyColumnOutOfRange {
114        /// Requested zero-based index.
115        index: usize,
116        /// Available columns.
117        columns: usize,
118    },
119    /// ABI v1 admits only scalar, non-floating key types.
120    #[error("shuffle key column {index} has unsupported partition type {data_type}")]
121    UnsupportedKeyType {
122        /// Requested zero-based index.
123        index: usize,
124        /// Rejected Arrow type.
125        data_type: arrow_schema::DataType,
126    },
127    /// Arrow row encoding or slicing failed.
128    #[error("shuffle Arrow routing: {0}")]
129    Arrow(#[from] arrow_schema::ArrowError),
130    /// Routing metadata must cover every input row exactly once.
131    #[error("shuffle route has {vnodes} vnode entries for {rows} rows")]
132    RowVnodeCardinality {
133        /// Input row count.
134        rows: usize,
135        /// Supplied route count.
136        vnodes: usize,
137    },
138    /// `UInt32` take indices cannot represent this input cardinality.
139    #[error("shuffle batch has {rows} rows; maximum supported is {}", u32::MAX)]
140    TooManyRows {
141        /// Input row count.
142        rows: usize,
143    },
144    /// The route references a vnode outside the pinned assignment.
145    #[error("shuffle vnode {vnode} is outside the {vnode_count}-vnode assignment")]
146    VnodeOutOfRange {
147        /// Invalid vnode.
148        vnode: u32,
149        /// Pinned assignment cardinality.
150        vnode_count: usize,
151    },
152    /// Formation/rebalance has not assigned this vnode yet.
153    #[error("shuffle vnode {vnode} is unassigned")]
154    UnassignedVnode {
155        /// Vnode awaiting an owner.
156        vnode: u32,
157    },
158    /// A single row cannot be split further and would always fail receiver admission.
159    #[error("shuffle row for vnode {vnode} occupies {bytes} decoded bytes; hard limit is {limit}")]
160    OversizedRow {
161        /// Row's routed vnode.
162        vnode: u32,
163        /// Decoded Arrow bytes.
164        bytes: usize,
165        /// Hard decoded limit.
166        limit: usize,
167    },
168}
169
170impl ShuffleRoutingError {
171    /// Whether retrying after assignment convergence can succeed without changing the input.
172    #[must_use]
173    pub const fn is_not_ready(&self) -> bool {
174        matches!(self, Self::UnassignedVnode { .. })
175    }
176}
177
178/// Hash each row with the engine's canonical Arrow-row and xxh3 encoding.
179///
180/// # Errors
181/// Returns a structural or Arrow encoding error instead of panicking on invalid resolved columns.
182pub 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        // Dictionary indices are an encoding detail. Hash the hydrated scalar
228        // value, but apply the same ABI gate recursively to that value type.
229        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        // Run-end encoding is also representation-level, but is excluded until
243        // equivalence with plain arrays is frozen by vectors.
244        arrow_schema::DataType::RunEndEncoded(_, _) => false,
245        data_type => !data_type.is_floating() && !data_type.is_nested(),
246    }
247}
248
249/// Build a complete local/remote plan from one caller-pinned assignment.
250///
251/// Every input row is assigned exactly once or the whole call fails. Remote batches carry vnode
252/// ownership out-of-band; user schemas, including a user field named `__laminar_vnode`, are never
253/// rewritten.
254///
255/// # Errors
256/// Returns before producing a plan for malformed metadata, unassigned/out-of-range vnodes, Arrow
257/// slicing failures, or a single row above the receiver's hard decoded-memory bound.
258pub 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], &registry.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            &registry.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            &registry.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], &registry.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}