1use std::{collections::VecDeque, fmt::Debug, pin::Pin, sync::Arc, task::Poll};
19
20use crate::{FlightData, FlightDescriptor, SchemaAsIpc, error::Result};
21
22use arrow_array::{Array, ArrayRef, RecordBatch, RecordBatchOptions, UnionArray};
23use arrow_ipc::writer::{DictionaryTracker, IpcDataGenerator, IpcWriteContext, IpcWriteOptions};
24
25use arrow_schema::{DataType, Field, FieldRef, Fields, Schema, SchemaRef, UnionMode};
26use bytes::Bytes;
27use futures::{Stream, StreamExt, ready, stream::BoxStream};
28
29#[derive(Debug)]
145pub struct FlightDataEncoderBuilder {
146 max_flight_data_size: usize,
149 options: IpcWriteOptions,
151 app_metadata: Bytes,
153 schema: Option<SchemaRef>,
155 descriptor: Option<FlightDescriptor>,
157 dictionary_handling: DictionaryHandling,
160}
161
162pub const GRPC_TARGET_MAX_FLIGHT_SIZE_BYTES: usize = 2097152;
167
168impl Default for FlightDataEncoderBuilder {
169 fn default() -> Self {
170 Self {
171 max_flight_data_size: GRPC_TARGET_MAX_FLIGHT_SIZE_BYTES,
172 options: IpcWriteOptions::default(),
173 app_metadata: Bytes::new(),
174 schema: None,
175 descriptor: None,
176 dictionary_handling: DictionaryHandling::Hydrate,
177 }
178 }
179}
180
181impl FlightDataEncoderBuilder {
182 pub fn new() -> Self {
184 Self::default()
185 }
186
187 pub fn with_max_flight_data_size(mut self, max_flight_data_size: usize) -> Self {
198 self.max_flight_data_size = max_flight_data_size;
199 self
200 }
201
202 pub fn with_dictionary_handling(mut self, dictionary_handling: DictionaryHandling) -> Self {
204 self.dictionary_handling = dictionary_handling;
205 self
206 }
207
208 pub fn with_metadata(mut self, app_metadata: Bytes) -> Self {
212 self.app_metadata = app_metadata;
213 self
214 }
215
216 pub fn with_options(mut self, options: IpcWriteOptions) -> Self {
218 self.options = options;
219 self
220 }
221
222 pub fn with_schema(mut self, schema: SchemaRef) -> Self {
227 self.schema = Some(schema);
228 self
229 }
230
231 pub fn with_flight_descriptor(mut self, descriptor: Option<FlightDescriptor>) -> Self {
233 self.descriptor = descriptor;
234 self
235 }
236
237 pub fn build<S>(self, input: S) -> FlightDataEncoder
242 where
243 S: Stream<Item = Result<RecordBatch>> + Send + 'static,
244 {
245 let Self {
246 max_flight_data_size,
247 options,
248 app_metadata,
249 schema,
250 descriptor,
251 dictionary_handling,
252 } = self;
253
254 FlightDataEncoder::new(
255 input.boxed(),
256 schema,
257 max_flight_data_size,
258 options,
259 app_metadata,
260 descriptor,
261 dictionary_handling,
262 )
263 }
264}
265
266pub struct FlightDataEncoder {
270 inner: BoxStream<'static, Result<RecordBatch>>,
272 schema: Option<SchemaRef>,
274 max_flight_data_size: usize,
277 encoder: FlightIpcEncoder,
279 app_metadata: Option<Bytes>,
281 queue: VecDeque<FlightData>,
283 done: bool,
285 descriptor: Option<FlightDescriptor>,
287 dictionary_handling: DictionaryHandling,
290}
291
292impl FlightDataEncoder {
293 fn new(
294 inner: BoxStream<'static, Result<RecordBatch>>,
295 schema: Option<SchemaRef>,
296 max_flight_data_size: usize,
297 options: IpcWriteOptions,
298 app_metadata: Bytes,
299 descriptor: Option<FlightDescriptor>,
300 dictionary_handling: DictionaryHandling,
301 ) -> Self {
302 let mut encoder = Self {
303 inner,
304 schema: None,
305 max_flight_data_size,
306 encoder: FlightIpcEncoder::new(
307 options,
308 dictionary_handling != DictionaryHandling::Resend,
309 ),
310 app_metadata: Some(app_metadata),
311 queue: VecDeque::new(),
312 done: false,
313 descriptor,
314 dictionary_handling,
315 };
316
317 if let Some(schema) = schema {
319 encoder.encode_schema(&schema);
320 }
321
322 encoder
323 }
324
325 pub fn known_schema(&self) -> Option<SchemaRef> {
328 self.schema.clone()
329 }
330
331 #[inline]
333 fn queue_message(&mut self, mut data: FlightData) {
334 if let Some(descriptor) = self.descriptor.take() {
335 data.flight_descriptor = Some(descriptor);
336 }
337 self.queue.push_back(data);
338 }
339
340 fn encode_schema(&mut self, schema: &SchemaRef) -> SchemaRef {
343 let send_dictionaries = self.dictionary_handling == DictionaryHandling::Resend;
346 let schema = Arc::new(prepare_schema_for_flight(
347 schema,
348 &mut self.encoder.dictionary_tracker,
349 send_dictionaries,
350 ));
351 let mut schema_flight_data = self.encoder.encode_schema(&schema);
352
353 if let Some(app_metadata) = self.app_metadata.take() {
355 schema_flight_data.app_metadata = app_metadata;
356 }
357 self.queue_message(schema_flight_data);
358 self.schema = Some(schema.clone());
360 schema
361 }
362
363 fn encode_batch(&mut self, batch: RecordBatch) -> Result<()> {
365 let schema = match &self.schema {
366 Some(schema) => schema.clone(),
367 None => self.encode_schema(batch.schema_ref()),
369 };
370
371 let batch = match self.dictionary_handling {
372 DictionaryHandling::Resend => batch,
373 DictionaryHandling::Hydrate => hydrate_dictionaries(&batch, schema)?,
374 };
375
376 let batches = split_batch_for_grpc_response(batch, self.max_flight_data_size);
377 let last = batches.len().saturating_sub(1); for (i, batch) in batches.into_iter().enumerate() {
379 self.encoder
380 .ipc_write_context
381 .set_reserve_scratch(i != last);
382 let (flight_dictionaries, flight_batch) = self.encoder.encode_batch(&batch)?;
383 for dict in flight_dictionaries {
384 self.queue_message(dict);
385 }
386 self.queue_message(flight_batch);
387 }
388
389 Ok(())
390 }
391}
392
393impl Stream for FlightDataEncoder {
394 type Item = Result<FlightData>;
395
396 fn poll_next(
397 mut self: Pin<&mut Self>,
398 cx: &mut std::task::Context<'_>,
399 ) -> Poll<Option<Self::Item>> {
400 loop {
401 if self.done && self.queue.is_empty() {
402 return Poll::Ready(None);
403 }
404
405 if let Some(data) = self.queue.pop_front() {
407 return Poll::Ready(Some(Ok(data)));
408 }
409
410 let batch = ready!(self.inner.poll_next_unpin(cx));
412
413 match batch {
414 None => {
415 self.done = true;
417 assert!(self.queue.is_empty());
419 return Poll::Ready(None);
420 }
421 Some(Err(e)) => {
422 self.done = true;
424 self.queue.clear();
425 return Poll::Ready(Some(Err(e)));
426 }
427 Some(Ok(batch)) => {
428 if let Err(e) = self.encode_batch(batch) {
430 self.done = true;
431 self.queue.clear();
432 return Poll::Ready(Some(Err(e)));
433 }
434 }
435 }
436 }
437 }
438}
439
440#[derive(Debug, PartialEq)]
469pub enum DictionaryHandling {
470 Hydrate,
478 Resend,
488}
489
490fn prepare_field_for_flight(
491 field: &FieldRef,
492 dictionary_tracker: &mut DictionaryTracker,
493 send_dictionaries: bool,
494) -> Field {
495 match field.data_type() {
496 DataType::List(inner) => Field::new_list(
497 field.name(),
498 prepare_field_for_flight(inner, dictionary_tracker, send_dictionaries),
499 field.is_nullable(),
500 )
501 .with_metadata(field.metadata().clone()),
502 DataType::LargeList(inner) => Field::new_list(
503 field.name(),
504 prepare_field_for_flight(inner, dictionary_tracker, send_dictionaries),
505 field.is_nullable(),
506 )
507 .with_metadata(field.metadata().clone()),
508 DataType::Struct(fields) => {
509 let new_fields: Vec<Field> = fields
510 .iter()
511 .map(|f| prepare_field_for_flight(f, dictionary_tracker, send_dictionaries))
512 .collect();
513 Field::new_struct(field.name(), new_fields, field.is_nullable())
514 .with_metadata(field.metadata().clone())
515 }
516 DataType::Union(fields, mode) => {
517 let (type_ids, new_fields): (Vec<i8>, Vec<Field>) = fields
518 .iter()
519 .map(|(type_id, f)| {
520 (
521 type_id,
522 prepare_field_for_flight(f, dictionary_tracker, send_dictionaries),
523 )
524 })
525 .unzip();
526
527 Field::new_union(field.name(), type_ids, new_fields, *mode)
528 }
529 DataType::Dictionary(_, value_type) => {
530 if !send_dictionaries {
531 let value_field = Field::new(
533 field.name(),
534 value_type.as_ref().clone(),
535 field.is_nullable(),
536 );
537 prepare_field_for_flight(
538 &Arc::new(value_field),
539 dictionary_tracker,
540 send_dictionaries,
541 )
542 .with_metadata(field.metadata().clone())
543 } else {
544 let value_field = Field::new("values", value_type.as_ref().clone(), true);
548 prepare_field_for_flight(
549 &Arc::new(value_field),
550 dictionary_tracker,
551 send_dictionaries,
552 );
553 dictionary_tracker.next_dict_id();
554 #[allow(deprecated)]
555 Field::new_dict(
556 field.name(),
557 field.data_type().clone(),
558 field.is_nullable(),
559 0,
560 field.dict_is_ordered().unwrap_or_default(),
561 )
562 .with_metadata(field.metadata().clone())
563 }
564 }
565 DataType::ListView(inner) | DataType::LargeListView(inner) => {
566 let prepared = prepare_field_for_flight(inner, dictionary_tracker, send_dictionaries);
567 Field::new(
568 field.name(),
569 match field.data_type() {
570 DataType::ListView(_) => DataType::ListView(Arc::new(prepared)),
571 _ => DataType::LargeListView(Arc::new(prepared)),
572 },
573 field.is_nullable(),
574 )
575 .with_metadata(field.metadata().clone())
576 }
577 DataType::FixedSizeList(inner, size) => Field::new(
578 field.name(),
579 DataType::FixedSizeList(
580 Arc::new(prepare_field_for_flight(
581 inner,
582 dictionary_tracker,
583 send_dictionaries,
584 )),
585 *size,
586 ),
587 field.is_nullable(),
588 )
589 .with_metadata(field.metadata().clone()),
590 DataType::RunEndEncoded(run_ends, values) => Field::new(
591 field.name(),
592 DataType::RunEndEncoded(
593 run_ends.clone(),
594 Arc::new(prepare_field_for_flight(
595 values,
596 dictionary_tracker,
597 send_dictionaries,
598 )),
599 ),
600 field.is_nullable(),
601 )
602 .with_metadata(field.metadata().clone()),
603 DataType::Map(inner, sorted) => Field::new(
604 field.name(),
605 DataType::Map(
606 prepare_field_for_flight(inner, dictionary_tracker, send_dictionaries).into(),
607 *sorted,
608 ),
609 field.is_nullable(),
610 )
611 .with_metadata(field.metadata().clone()),
612 DataType::Null
613 | DataType::Boolean
614 | DataType::Int8
615 | DataType::Int16
616 | DataType::Int32
617 | DataType::Int64
618 | DataType::UInt8
619 | DataType::UInt16
620 | DataType::UInt32
621 | DataType::UInt64
622 | DataType::Float16
623 | DataType::Float32
624 | DataType::Float64
625 | DataType::Timestamp(_, _)
626 | DataType::Date32
627 | DataType::Date64
628 | DataType::Time32(_)
629 | DataType::Time64(_)
630 | DataType::Duration(_)
631 | DataType::Interval(_)
632 | DataType::Binary
633 | DataType::FixedSizeBinary(_)
634 | DataType::LargeBinary
635 | DataType::BinaryView
636 | DataType::Utf8
637 | DataType::LargeUtf8
638 | DataType::Utf8View
639 | DataType::Decimal32(_, _)
640 | DataType::Decimal64(_, _)
641 | DataType::Decimal128(_, _)
642 | DataType::Decimal256(_, _) => field.as_ref().clone(),
643 }
644}
645
646fn prepare_schema_for_flight(
652 schema: &Schema,
653 dictionary_tracker: &mut DictionaryTracker,
654 send_dictionaries: bool,
655) -> Schema {
656 let fields: Fields = schema
657 .fields()
658 .iter()
659 .map(|field| prepare_field_for_flight(field, dictionary_tracker, send_dictionaries))
660 .collect();
661
662 Schema::new(fields).with_metadata(schema.metadata().clone())
663}
664
665fn split_batch_for_grpc_response(
672 batch: RecordBatch,
673 max_flight_data_size: usize,
674) -> Vec<RecordBatch> {
675 let size = batch
676 .columns()
677 .iter()
678 .map(|col| col.get_buffer_memory_size())
679 .sum::<usize>();
680
681 let n_batches =
682 (size / max_flight_data_size + usize::from(size % max_flight_data_size != 0)).max(1);
683 let num_rows = batch.num_rows();
684 let rows_per_batch = (num_rows / n_batches).max(1);
685 let mut offset = 0;
686 let mut batches = Vec::with_capacity(n_batches);
687
688 while offset < num_rows {
689 let length = rows_per_batch.min(num_rows - offset);
690 batches.push(batch.slice(offset, length));
691 offset += length;
692 }
693
694 batches
695}
696
697struct FlightIpcEncoder {
704 options: IpcWriteOptions,
705 data_gen: IpcDataGenerator,
706 dictionary_tracker: DictionaryTracker,
707 ipc_write_context: IpcWriteContext,
708}
709
710impl FlightIpcEncoder {
711 fn new(options: IpcWriteOptions, error_on_replacement: bool) -> Self {
712 Self {
713 options,
714 data_gen: IpcDataGenerator::default(),
715 dictionary_tracker: DictionaryTracker::new(error_on_replacement),
716 ipc_write_context: IpcWriteContext::default(),
717 }
718 }
719
720 fn encode_schema(&self, schema: &Schema) -> FlightData {
722 SchemaAsIpc::new(schema, &self.options).into()
723 }
724
725 fn encode_batch(
728 &mut self,
729 batch: &RecordBatch,
730 ) -> Result<(impl Iterator<Item = FlightData> + use<>, FlightData)> {
731 let (encoded_dictionaries, encoded_batch) = self.data_gen.encode(
732 batch,
733 &mut self.dictionary_tracker,
734 &self.options,
735 &mut self.ipc_write_context,
736 )?;
737
738 let flight_dictionaries = encoded_dictionaries.into_iter().map(|e| e.into());
739 let flight_batch = encoded_batch.into();
740
741 Ok((flight_dictionaries, flight_batch))
742 }
743}
744
745fn hydrate_dictionaries(batch: &RecordBatch, schema: SchemaRef) -> Result<RecordBatch> {
748 let columns = schema
749 .fields()
750 .iter()
751 .zip(batch.columns())
752 .map(|(field, c)| hydrate_dictionary(c, field.data_type()))
753 .collect::<Result<Vec<_>>>()?;
754
755 let options = RecordBatchOptions::new().with_row_count(Some(batch.num_rows()));
756
757 Ok(RecordBatch::try_new_with_options(
758 schema, columns, &options,
759 )?)
760}
761
762fn hydrate_dictionary(array: &ArrayRef, data_type: &DataType) -> Result<ArrayRef> {
764 let arr = match (array.data_type(), data_type) {
765 (DataType::Union(_, UnionMode::Sparse), DataType::Union(fields, UnionMode::Sparse)) => {
766 let union_arr = array.as_any().downcast_ref::<UnionArray>().unwrap();
767
768 Arc::new(UnionArray::try_new(
769 fields.clone(),
770 union_arr.type_ids().clone(),
771 None,
772 fields
773 .iter()
774 .map(|(type_id, field)| {
775 Ok(arrow_cast::cast(
776 union_arr.child(type_id),
777 field.data_type(),
778 )?)
779 })
780 .collect::<Result<Vec<_>>>()?,
781 )?)
782 }
783 (_, data_type) => arrow_cast::cast(array, data_type)?,
784 };
785 Ok(arr)
786}
787
788#[cfg(test)]
789mod tests {
790 use crate::decode::{DecodedPayload, FlightDataDecoder};
791 use arrow_array::builder::{
792 FixedSizeListBuilder, GenericByteDictionaryBuilder, GenericListViewBuilder, ListBuilder,
793 StringDictionaryBuilder, StructBuilder,
794 };
795 use arrow_array::*;
796 use arrow_array::{cast::downcast_array, types::*};
797 use arrow_buffer::ScalarBuffer;
798 use arrow_cast::pretty::pretty_format_batches;
799 use arrow_ipc::{CompressionType, MetadataVersion};
800 use arrow_schema::{UnionFields, UnionMode};
801 use builder::MapBuilder;
802 use std::collections::HashMap;
803
804 use super::*;
805
806 #[test]
807 fn test_encode_flight_data() {
810 let options = IpcWriteOptions::try_new(8, false, MetadataVersion::V5).unwrap();
812 let c1 = UInt32Array::from(vec![1, 2, 3, 4, 5, 6]);
813
814 let batch = RecordBatch::try_from_iter(vec![("a", Arc::new(c1) as ArrayRef)])
815 .expect("cannot create record batch");
816 let schema = batch.schema_ref();
817
818 let (_, baseline_flight_batch) = make_flight_data(&batch, &options);
819
820 let big_batch = batch.slice(0, batch.num_rows() - 1);
821 let optimized_big_batch =
822 hydrate_dictionaries(&big_batch, Arc::clone(schema)).expect("failed to optimize");
823 let (_, optimized_big_flight_batch) = make_flight_data(&optimized_big_batch, &options);
824
825 assert_eq!(
826 baseline_flight_batch.data_body.len(),
827 optimized_big_flight_batch.data_body.len()
828 );
829
830 let small_batch = batch.slice(0, 1);
831 let optimized_small_batch =
832 hydrate_dictionaries(&small_batch, Arc::clone(schema)).expect("failed to optimize");
833 let (_, optimized_small_flight_batch) = make_flight_data(&optimized_small_batch, &options);
834
835 assert!(
836 baseline_flight_batch.data_body.len() > optimized_small_flight_batch.data_body.len()
837 );
838 }
839
840 #[tokio::test]
841 async fn test_dictionary_hydration() {
842 let arr1: DictionaryArray<UInt16Type> = vec!["a", "a", "b"].into_iter().collect();
843 let arr2: DictionaryArray<UInt16Type> = vec!["c", "c", "d"].into_iter().collect();
844
845 let schema = Arc::new(Schema::new(vec![Field::new_dictionary(
846 "dict",
847 DataType::UInt16,
848 DataType::Utf8,
849 false,
850 )]));
851 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
852 let batch2 = RecordBatch::try_new(schema, vec![Arc::new(arr2)]).unwrap();
853
854 let stream = futures::stream::iter(vec![Ok(batch1), Ok(batch2)]);
855
856 let encoder = FlightDataEncoderBuilder::default().build(stream);
857 let mut decoder = FlightDataDecoder::new(encoder);
858 let expected_schema = Schema::new(vec![Field::new("dict", DataType::Utf8, false)]);
859 let expected_schema = Arc::new(expected_schema);
860 let mut expected_arrays = vec![
861 StringArray::from(vec!["a", "a", "b"]),
862 StringArray::from(vec!["c", "c", "d"]),
863 ]
864 .into_iter();
865 while let Some(decoded) = decoder.next().await {
866 let decoded = decoded.unwrap();
867 match decoded.payload {
868 DecodedPayload::None => {}
869 DecodedPayload::Schema(s) => assert_eq!(s, expected_schema),
870 DecodedPayload::RecordBatch(b) => {
871 assert_eq!(b.schema(), expected_schema);
872 let expected_array = expected_arrays.next().unwrap();
873 let actual_array = b.column_by_name("dict").unwrap();
874 let actual_array = downcast_array::<StringArray>(actual_array);
875
876 assert_eq!(actual_array, expected_array);
877 }
878 }
879 }
880 }
881
882 #[tokio::test]
883 async fn test_dictionary_resend() {
884 let arr1: DictionaryArray<UInt16Type> = vec!["a", "a", "b"].into_iter().collect();
885 let arr2: DictionaryArray<UInt16Type> = vec!["c", "c", "d"].into_iter().collect();
886
887 let schema = Arc::new(Schema::new(vec![Field::new_dictionary(
888 "dict",
889 DataType::UInt16,
890 DataType::Utf8,
891 false,
892 )]));
893 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
894 let batch2 = RecordBatch::try_new(schema, vec![Arc::new(arr2)]).unwrap();
895
896 verify_flight_round_trip(vec![batch1, batch2]).await;
897 }
898
899 #[tokio::test]
900 async fn test_compression_round_trip() {
901 let ints = Int32Array::from_iter_values((0..1024).map(|i| i % 8));
905 let strings = StringArray::from_iter_values((0..1024).map(|i| format!("value-{}", i % 8)));
906 let batch = RecordBatch::try_from_iter(vec![
907 ("ints", Arc::new(ints) as ArrayRef),
908 ("strings", Arc::new(strings) as ArrayRef),
909 ])
910 .unwrap();
911
912 for compression in [CompressionType::LZ4_FRAME, CompressionType::ZSTD] {
913 let options = IpcWriteOptions::default()
914 .try_with_compression(Some(compression))
915 .unwrap();
916 verify_flight_round_trip_with_options(vec![batch.clone()], options).await;
917 }
918 }
919
920 #[tokio::test]
921 async fn test_dictionary_hydration_known_schema() {
922 let arr1: DictionaryArray<UInt16Type> = vec!["a", "a", "b"].into_iter().collect();
923 let arr2: DictionaryArray<UInt16Type> = vec!["c", "c", "d"].into_iter().collect();
924
925 let schema = Arc::new(Schema::new(vec![Field::new_dictionary(
926 "dict",
927 DataType::UInt16,
928 DataType::Utf8,
929 false,
930 )]));
931 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
932 let batch2 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr2)]).unwrap();
933
934 let stream = futures::stream::iter(vec![Ok(batch1), Ok(batch2)]);
935
936 let encoder = FlightDataEncoderBuilder::default()
937 .with_schema(schema)
938 .build(stream);
939 let expected_schema =
940 Arc::new(Schema::new(vec![Field::new("dict", DataType::Utf8, false)]));
941 assert_eq!(Some(expected_schema), encoder.known_schema())
942 }
943
944 #[tokio::test]
945 async fn test_dictionary_resend_known_schema() {
946 let arr1: DictionaryArray<UInt16Type> = vec!["a", "a", "b"].into_iter().collect();
947 let arr2: DictionaryArray<UInt16Type> = vec!["c", "c", "d"].into_iter().collect();
948
949 let schema = Arc::new(Schema::new(vec![Field::new_dictionary(
950 "dict",
951 DataType::UInt16,
952 DataType::Utf8,
953 false,
954 )]));
955 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
956 let batch2 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr2)]).unwrap();
957
958 let stream = futures::stream::iter(vec![Ok(batch1), Ok(batch2)]);
959
960 let encoder = FlightDataEncoderBuilder::default()
961 .with_dictionary_handling(DictionaryHandling::Resend)
962 .with_schema(schema.clone())
963 .build(stream);
964 assert_eq!(Some(schema), encoder.known_schema())
965 }
966
967 #[tokio::test]
968 async fn test_multiple_dictionaries_resend() {
969 let schema = Arc::new(Schema::new(vec![
971 Field::new_dictionary("dict_1", DataType::UInt16, DataType::Utf8, false),
972 Field::new_dictionary("dict_2", DataType::UInt16, DataType::Utf8, false),
973 ]));
974
975 let arr_one_1: Arc<DictionaryArray<UInt16Type>> =
976 Arc::new(vec!["a", "a", "b"].into_iter().collect());
977 let arr_one_2: Arc<DictionaryArray<UInt16Type>> =
978 Arc::new(vec!["c", "c", "d"].into_iter().collect());
979 let arr_two_1: Arc<DictionaryArray<UInt16Type>> =
980 Arc::new(vec!["b", "a", "c"].into_iter().collect());
981 let arr_two_2: Arc<DictionaryArray<UInt16Type>> =
982 Arc::new(vec!["k", "d", "e"].into_iter().collect());
983 let batch1 =
984 RecordBatch::try_new(schema.clone(), vec![arr_one_1.clone(), arr_one_2.clone()])
985 .unwrap();
986 let batch2 =
987 RecordBatch::try_new(schema.clone(), vec![arr_two_1.clone(), arr_two_2.clone()])
988 .unwrap();
989
990 verify_flight_round_trip(vec![batch1, batch2]).await;
991 }
992
993 #[tokio::test]
994 async fn test_dictionary_list_hydration() {
995 let mut builder = ListBuilder::new(StringDictionaryBuilder::<UInt16Type>::new());
996
997 builder.append_value(vec![Some("a"), None, Some("b")]);
998
999 let arr1 = builder.finish();
1000
1001 builder.append_value(vec![Some("c"), None, Some("d")]);
1002
1003 let arr2 = builder.finish();
1004
1005 let schema = Arc::new(Schema::new(vec![Field::new_list(
1006 "dict_list",
1007 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1008 true,
1009 )]));
1010
1011 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1012 let batch2 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr2)]).unwrap();
1013
1014 let stream = futures::stream::iter(vec![Ok(batch1), Ok(batch2)]);
1015
1016 let encoder = FlightDataEncoderBuilder::default().build(stream);
1017
1018 let mut decoder = FlightDataDecoder::new(encoder);
1019 let expected_schema = Schema::new(vec![Field::new_list(
1020 "dict_list",
1021 Field::new_list_field(DataType::Utf8, true),
1022 true,
1023 )]);
1024
1025 let expected_schema = Arc::new(expected_schema);
1026
1027 let mut expected_arrays = vec![
1028 StringArray::from_iter(vec![Some("a"), None, Some("b")]),
1029 StringArray::from_iter(vec![Some("c"), None, Some("d")]),
1030 ]
1031 .into_iter();
1032
1033 while let Some(decoded) = decoder.next().await {
1034 let decoded = decoded.unwrap();
1035 match decoded.payload {
1036 DecodedPayload::None => {}
1037 DecodedPayload::Schema(s) => assert_eq!(s, expected_schema),
1038 DecodedPayload::RecordBatch(b) => {
1039 assert_eq!(b.schema(), expected_schema);
1040 let expected_array = expected_arrays.next().unwrap();
1041 let list_array =
1042 downcast_array::<ListArray>(b.column_by_name("dict_list").unwrap());
1043 let elem_array = downcast_array::<StringArray>(list_array.value(0).as_ref());
1044
1045 assert_eq!(elem_array, expected_array);
1046 }
1047 }
1048 }
1049 }
1050
1051 #[tokio::test]
1052 async fn test_dictionary_list_resend() {
1053 let mut builder = ListBuilder::new(StringDictionaryBuilder::<UInt16Type>::new());
1054
1055 builder.append_value(vec![Some("a"), None, Some("b")]);
1056
1057 let arr1 = builder.finish();
1058
1059 builder.append_value(vec![Some("c"), None, Some("d")]);
1060
1061 let arr2 = builder.finish();
1062
1063 let schema = Arc::new(Schema::new(vec![Field::new_list(
1064 "dict_list",
1065 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1066 true,
1067 )]));
1068
1069 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1070 let batch2 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr2)]).unwrap();
1071
1072 verify_flight_round_trip(vec![batch1, batch2]).await;
1073 }
1074
1075 #[tokio::test]
1076 async fn test_dictionary_struct_hydration() {
1077 let struct_fields = vec![Field::new_list(
1078 "dict_list",
1079 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1080 true,
1081 )];
1082
1083 let mut struct_builder = StructBuilder::new(
1084 struct_fields.clone(),
1085 vec![Box::new(builder::ListBuilder::new(
1086 StringDictionaryBuilder::<UInt16Type>::new(),
1087 ))],
1088 );
1089
1090 struct_builder
1091 .field_builder::<ListBuilder<GenericByteDictionaryBuilder<UInt16Type,GenericStringType<i32>>>>(0)
1092 .unwrap()
1093 .append_value(vec![Some("a"), None, Some("b")]);
1094
1095 struct_builder.append(true);
1096
1097 let arr1 = struct_builder.finish();
1098
1099 struct_builder
1100 .field_builder::<ListBuilder<GenericByteDictionaryBuilder<UInt16Type,GenericStringType<i32>>>>(0)
1101 .unwrap()
1102 .append_value(vec![Some("c"), None, Some("d")]);
1103 struct_builder.append(true);
1104
1105 let arr2 = struct_builder.finish();
1106
1107 let schema = Arc::new(Schema::new(vec![Field::new_struct(
1108 "struct",
1109 struct_fields,
1110 true,
1111 )]));
1112
1113 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1114 let batch2 = RecordBatch::try_new(schema, vec![Arc::new(arr2)]).unwrap();
1115
1116 let stream = futures::stream::iter(vec![Ok(batch1), Ok(batch2)]);
1117
1118 let encoder = FlightDataEncoderBuilder::default().build(stream);
1119
1120 let mut decoder = FlightDataDecoder::new(encoder);
1121 let expected_schema = Schema::new(vec![Field::new_struct(
1122 "struct",
1123 vec![Field::new_list(
1124 "dict_list",
1125 Field::new_list_field(DataType::Utf8, true),
1126 true,
1127 )],
1128 true,
1129 )]);
1130
1131 let expected_schema = Arc::new(expected_schema);
1132
1133 let mut expected_arrays = vec![
1134 StringArray::from_iter(vec![Some("a"), None, Some("b")]),
1135 StringArray::from_iter(vec![Some("c"), None, Some("d")]),
1136 ]
1137 .into_iter();
1138
1139 while let Some(decoded) = decoder.next().await {
1140 let decoded = decoded.unwrap();
1141 match decoded.payload {
1142 DecodedPayload::None => {}
1143 DecodedPayload::Schema(s) => assert_eq!(s, expected_schema),
1144 DecodedPayload::RecordBatch(b) => {
1145 assert_eq!(b.schema(), expected_schema);
1146 let expected_array = expected_arrays.next().unwrap();
1147 let struct_array =
1148 downcast_array::<StructArray>(b.column_by_name("struct").unwrap());
1149 let list_array = downcast_array::<ListArray>(struct_array.column(0));
1150
1151 let elem_array = downcast_array::<StringArray>(list_array.value(0).as_ref());
1152
1153 assert_eq!(elem_array, expected_array);
1154 }
1155 }
1156 }
1157 }
1158
1159 #[tokio::test]
1160 async fn test_dictionary_struct_resend() {
1161 let struct_fields = vec![Field::new_list(
1162 "dict_list",
1163 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1164 true,
1165 )];
1166
1167 let mut struct_builder = StructBuilder::new(
1168 struct_fields.clone(),
1169 vec![Box::new(builder::ListBuilder::new(
1170 StringDictionaryBuilder::<UInt16Type>::new(),
1171 ))],
1172 );
1173
1174 struct_builder.field_builder::<ListBuilder<GenericByteDictionaryBuilder<UInt16Type,GenericStringType<i32>>>>(0)
1175 .unwrap()
1176 .append_value(vec![Some("a"), None, Some("b")]);
1177 struct_builder.append(true);
1178
1179 let arr1 = struct_builder.finish();
1180
1181 struct_builder.field_builder::<ListBuilder<GenericByteDictionaryBuilder<UInt16Type,GenericStringType<i32>>>>(0)
1182 .unwrap()
1183 .append_value(vec![Some("c"), None, Some("d")]);
1184 struct_builder.append(true);
1185
1186 let arr2 = struct_builder.finish();
1187
1188 let schema = Arc::new(Schema::new(vec![Field::new_struct(
1189 "struct",
1190 struct_fields,
1191 true,
1192 )]));
1193
1194 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1195 let batch2 = RecordBatch::try_new(schema, vec![Arc::new(arr2)]).unwrap();
1196
1197 verify_flight_round_trip(vec![batch1, batch2]).await;
1198 }
1199
1200 #[tokio::test]
1201 async fn test_dictionary_union_hydration() {
1202 let struct_fields = vec![Field::new_list(
1203 "dict_list",
1204 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1205 true,
1206 )];
1207
1208 let union_fields = [
1209 (
1210 0,
1211 Arc::new(Field::new_list(
1212 "dict_list",
1213 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1214 true,
1215 )),
1216 ),
1217 (
1218 1,
1219 Arc::new(Field::new_struct("struct", struct_fields.clone(), true)),
1220 ),
1221 (2, Arc::new(Field::new("string", DataType::Utf8, true))),
1222 ]
1223 .into_iter()
1224 .collect::<UnionFields>();
1225
1226 let struct_fields = vec![Field::new_list(
1227 "dict_list",
1228 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1229 true,
1230 )];
1231
1232 let mut builder = builder::ListBuilder::new(StringDictionaryBuilder::<UInt16Type>::new());
1233
1234 builder.append_value(vec![Some("a"), None, Some("b")]);
1235
1236 let arr1 = builder.finish();
1237
1238 let type_id_buffer = [0].into_iter().collect::<ScalarBuffer<i8>>();
1239 let arr1 = UnionArray::try_new(
1240 union_fields.clone(),
1241 type_id_buffer,
1242 None,
1243 vec![
1244 Arc::new(arr1) as Arc<dyn Array>,
1245 new_null_array(union_fields.iter().nth(1).unwrap().1.data_type(), 1),
1246 new_null_array(union_fields.iter().nth(2).unwrap().1.data_type(), 1),
1247 ],
1248 )
1249 .unwrap();
1250
1251 builder.append_value(vec![Some("c"), None, Some("d")]);
1252
1253 let arr2 = Arc::new(builder.finish());
1254 let arr2 = StructArray::new(struct_fields.clone().into(), vec![arr2], None);
1255
1256 let type_id_buffer = [1].into_iter().collect::<ScalarBuffer<i8>>();
1257 let arr2 = UnionArray::try_new(
1258 union_fields.clone(),
1259 type_id_buffer,
1260 None,
1261 vec![
1262 new_null_array(union_fields.iter().next().unwrap().1.data_type(), 1),
1263 Arc::new(arr2),
1264 new_null_array(union_fields.iter().nth(2).unwrap().1.data_type(), 1),
1265 ],
1266 )
1267 .unwrap();
1268
1269 let type_id_buffer = [2].into_iter().collect::<ScalarBuffer<i8>>();
1270 let arr3 = UnionArray::try_new(
1271 union_fields.clone(),
1272 type_id_buffer,
1273 None,
1274 vec![
1275 new_null_array(union_fields.iter().next().unwrap().1.data_type(), 1),
1276 new_null_array(union_fields.iter().nth(1).unwrap().1.data_type(), 1),
1277 Arc::new(StringArray::from(vec!["e"])),
1278 ],
1279 )
1280 .unwrap();
1281
1282 let (type_ids, union_fields): (Vec<_>, Vec<_>) = union_fields
1283 .iter()
1284 .map(|(type_id, field_ref)| (type_id, (*Arc::clone(field_ref)).clone()))
1285 .unzip();
1286 let schema = Arc::new(Schema::new(vec![Field::new_union(
1287 "union",
1288 type_ids.clone(),
1289 union_fields.clone(),
1290 UnionMode::Sparse,
1291 )]));
1292
1293 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1294 let batch2 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr2)]).unwrap();
1295 let batch3 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr3)]).unwrap();
1296
1297 let stream = futures::stream::iter(vec![Ok(batch1), Ok(batch2), Ok(batch3)]);
1298
1299 let encoder = FlightDataEncoderBuilder::default().build(stream);
1300
1301 let mut decoder = FlightDataDecoder::new(encoder);
1302
1303 let hydrated_struct_fields = vec![Field::new_list(
1304 "dict_list",
1305 Field::new_list_field(DataType::Utf8, true),
1306 true,
1307 )];
1308
1309 let hydrated_union_fields = vec![
1310 Field::new_list(
1311 "dict_list",
1312 Field::new_list_field(DataType::Utf8, true),
1313 true,
1314 ),
1315 Field::new_struct("struct", hydrated_struct_fields.clone(), true),
1316 Field::new("string", DataType::Utf8, true),
1317 ];
1318
1319 let expected_schema = Schema::new(vec![Field::new_union(
1320 "union",
1321 type_ids.clone(),
1322 hydrated_union_fields,
1323 UnionMode::Sparse,
1324 )]);
1325
1326 let expected_schema = Arc::new(expected_schema);
1327
1328 let mut expected_arrays = vec![
1329 StringArray::from_iter(vec![Some("a"), None, Some("b")]),
1330 StringArray::from_iter(vec![Some("c"), None, Some("d")]),
1331 StringArray::from(vec!["e"]),
1332 ]
1333 .into_iter();
1334
1335 let mut batch = 0;
1336 while let Some(decoded) = decoder.next().await {
1337 let decoded = decoded.unwrap();
1338 match decoded.payload {
1339 DecodedPayload::None => {}
1340 DecodedPayload::Schema(s) => assert_eq!(s, expected_schema),
1341 DecodedPayload::RecordBatch(b) => {
1342 assert_eq!(b.schema(), expected_schema);
1343 let expected_array = expected_arrays.next().unwrap();
1344 let union_arr =
1345 downcast_array::<UnionArray>(b.column_by_name("union").unwrap());
1346
1347 let elem_array = match batch {
1348 0 => {
1349 let list_array = downcast_array::<ListArray>(union_arr.child(0));
1350 downcast_array::<StringArray>(list_array.value(0).as_ref())
1351 }
1352 1 => {
1353 let struct_array = downcast_array::<StructArray>(union_arr.child(1));
1354 let list_array = downcast_array::<ListArray>(struct_array.column(0));
1355
1356 downcast_array::<StringArray>(list_array.value(0).as_ref())
1357 }
1358 _ => downcast_array::<StringArray>(union_arr.child(2)),
1359 };
1360
1361 batch += 1;
1362
1363 assert_eq!(elem_array, expected_array);
1364 }
1365 }
1366 }
1367 }
1368
1369 #[tokio::test]
1370 async fn test_dictionary_union_resend() {
1371 let struct_fields = vec![Field::new_list(
1372 "dict_list",
1373 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1374 true,
1375 )];
1376
1377 let union_fields = [
1378 (
1379 0,
1380 Arc::new(Field::new_list(
1381 "dict_list",
1382 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1383 true,
1384 )),
1385 ),
1386 (
1387 1,
1388 Arc::new(Field::new_struct("struct", struct_fields.clone(), true)),
1389 ),
1390 (2, Arc::new(Field::new("string", DataType::Utf8, true))),
1391 ]
1392 .into_iter()
1393 .collect::<UnionFields>();
1394
1395 let mut field_types = union_fields.iter().map(|(_, field)| field.data_type());
1396 let dict_list_ty = field_types.next().unwrap();
1397 let struct_ty = field_types.next().unwrap();
1398 let string_ty = field_types.next().unwrap();
1399
1400 let struct_fields = vec![Field::new_list(
1401 "dict_list",
1402 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1403 true,
1404 )];
1405
1406 let mut builder = builder::ListBuilder::new(StringDictionaryBuilder::<UInt16Type>::new());
1407
1408 builder.append_value(vec![Some("a"), None, Some("b")]);
1409
1410 let arr1 = builder.finish();
1411
1412 let type_id_buffer = [0].into_iter().collect::<ScalarBuffer<i8>>();
1413 let arr1 = UnionArray::try_new(
1414 union_fields.clone(),
1415 type_id_buffer,
1416 None,
1417 vec![
1418 Arc::new(arr1),
1419 new_null_array(struct_ty, 1),
1420 new_null_array(string_ty, 1),
1421 ],
1422 )
1423 .unwrap();
1424
1425 builder.append_value(vec![Some("c"), None, Some("d")]);
1426
1427 let arr2 = Arc::new(builder.finish());
1428 let arr2 = StructArray::new(struct_fields.clone().into(), vec![arr2], None);
1429
1430 let type_id_buffer = [1].into_iter().collect::<ScalarBuffer<i8>>();
1431 let arr2 = UnionArray::try_new(
1432 union_fields.clone(),
1433 type_id_buffer,
1434 None,
1435 vec![
1436 new_null_array(dict_list_ty, 1),
1437 Arc::new(arr2),
1438 new_null_array(string_ty, 1),
1439 ],
1440 )
1441 .unwrap();
1442
1443 let type_id_buffer = [2].into_iter().collect::<ScalarBuffer<i8>>();
1444 let arr3 = UnionArray::try_new(
1445 union_fields.clone(),
1446 type_id_buffer,
1447 None,
1448 vec![
1449 new_null_array(dict_list_ty, 1),
1450 new_null_array(struct_ty, 1),
1451 Arc::new(StringArray::from(vec!["e"])),
1452 ],
1453 )
1454 .unwrap();
1455
1456 let (type_ids, union_fields): (Vec<_>, Vec<_>) = union_fields
1457 .iter()
1458 .map(|(type_id, field_ref)| (type_id, (*Arc::clone(field_ref)).clone()))
1459 .unzip();
1460 let schema = Arc::new(Schema::new(vec![Field::new_union(
1461 "union",
1462 type_ids.clone(),
1463 union_fields.clone(),
1464 UnionMode::Sparse,
1465 )]));
1466
1467 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1468 let batch2 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr2)]).unwrap();
1469 let batch3 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr3)]).unwrap();
1470
1471 verify_flight_round_trip(vec![batch1, batch2, batch3]).await;
1472 }
1473
1474 #[tokio::test]
1475 async fn test_dictionary_map_hydration() {
1476 let mut builder = MapBuilder::new(
1477 None,
1478 StringDictionaryBuilder::<UInt16Type>::new(),
1479 StringDictionaryBuilder::<UInt16Type>::new(),
1480 );
1481
1482 builder.keys().append_value("k1");
1484 builder.values().append_value("a");
1485 builder.keys().append_value("k2");
1486 builder.values().append_null();
1487 builder.keys().append_value("k3");
1488 builder.values().append_value("b");
1489 builder.append(true).unwrap();
1490
1491 let arr1 = builder.finish();
1492
1493 builder.keys().append_value("k1");
1495 builder.values().append_value("c");
1496 builder.keys().append_value("k2");
1497 builder.values().append_null();
1498 builder.keys().append_value("k3");
1499 builder.values().append_value("d");
1500 builder.append(true).unwrap();
1501
1502 let arr2 = builder.finish();
1503
1504 let schema = Arc::new(Schema::new(vec![Field::new_map(
1505 "dict_map",
1506 Field::MAP_ENTRIES_FIELD_DEFAULT_NAME,
1507 Field::new_dictionary(
1508 Field::MAP_KEY_FIELD_DEFAULT_NAME,
1509 DataType::UInt16,
1510 DataType::Utf8,
1511 false,
1512 ),
1513 Field::new_dictionary(
1514 Field::MAP_VALUE_FIELD_DEFAULT_NAME,
1515 DataType::UInt16,
1516 DataType::Utf8,
1517 true,
1518 ),
1519 false,
1520 false,
1521 )]));
1522
1523 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1524 let batch2 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr2)]).unwrap();
1525
1526 let stream = futures::stream::iter(vec![Ok(batch1), Ok(batch2)]);
1527
1528 let encoder = FlightDataEncoderBuilder::default().build(stream);
1529
1530 let mut decoder = FlightDataDecoder::new(encoder);
1531 let expected_schema = Schema::new(vec![Field::new_map(
1532 "dict_map",
1533 Field::MAP_ENTRIES_FIELD_DEFAULT_NAME,
1534 Field::new(Field::MAP_KEY_FIELD_DEFAULT_NAME, DataType::Utf8, false),
1535 Field::new(Field::MAP_VALUE_FIELD_DEFAULT_NAME, DataType::Utf8, true),
1536 false,
1537 false,
1538 )]);
1539
1540 let expected_schema = Arc::new(expected_schema);
1541
1542 let arr1 = MapArray::from_vec_of_maps::<StringArray, StringArray, _, _>(
1544 vec![Some(vec![
1545 ("k1", Some("a")),
1546 ("k2", None),
1547 ("k3", Some("b")),
1548 ])],
1549 false,
1550 );
1551
1552 let arr2 = MapArray::from_vec_of_maps::<StringArray, StringArray, _, _>(
1553 vec![Some(vec![
1554 ("k1", Some("c")),
1555 ("k2", None),
1556 ("k3", Some("d")),
1557 ])],
1558 false,
1559 );
1560
1561 let mut expected_arrays = vec![arr1, arr2].into_iter();
1562
1563 while let Some(decoded) = decoder.next().await {
1564 let decoded = decoded.unwrap();
1565 match decoded.payload {
1566 DecodedPayload::None => {}
1567 DecodedPayload::Schema(s) => assert_eq!(s, expected_schema),
1568 DecodedPayload::RecordBatch(b) => {
1569 assert_eq!(b.schema(), expected_schema);
1570 let expected_array = expected_arrays.next().unwrap();
1571 let map_array =
1572 downcast_array::<MapArray>(b.column_by_name("dict_map").unwrap());
1573
1574 assert_eq!(map_array, expected_array);
1575 }
1576 }
1577 }
1578 }
1579
1580 #[tokio::test]
1581 async fn test_dictionary_map_resend() {
1582 let mut builder = MapBuilder::new(
1583 None,
1584 StringDictionaryBuilder::<UInt16Type>::new(),
1585 StringDictionaryBuilder::<UInt16Type>::new(),
1586 );
1587
1588 builder.keys().append_value("k1");
1590 builder.values().append_value("a");
1591 builder.keys().append_value("k2");
1592 builder.values().append_null();
1593 builder.keys().append_value("k3");
1594 builder.values().append_value("b");
1595 builder.append(true).unwrap();
1596
1597 let arr1 = builder.finish();
1598
1599 builder.keys().append_value("k1");
1601 builder.values().append_value("c");
1602 builder.keys().append_value("k2");
1603 builder.values().append_null();
1604 builder.keys().append_value("k3");
1605 builder.values().append_value("d");
1606 builder.append(true).unwrap();
1607
1608 let arr2 = builder.finish();
1609
1610 let schema = Arc::new(Schema::new(vec![Field::new_map(
1611 "dict_map",
1612 Field::MAP_ENTRIES_FIELD_DEFAULT_NAME,
1613 Field::new_dictionary(
1614 Field::MAP_KEY_FIELD_DEFAULT_NAME,
1615 DataType::UInt16,
1616 DataType::Utf8,
1617 false,
1618 ),
1619 Field::new_dictionary(
1620 Field::MAP_VALUE_FIELD_DEFAULT_NAME,
1621 DataType::UInt16,
1622 DataType::Utf8,
1623 true,
1624 ),
1625 false,
1626 false,
1627 )]));
1628
1629 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1630 let batch2 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr2)]).unwrap();
1631
1632 verify_flight_round_trip(vec![batch1, batch2]).await;
1633 }
1634
1635 #[tokio::test]
1636 async fn test_dictionary_ree_resend() {
1637 let dict_values1 = vec![Some("a"), None, Some("b")]
1638 .into_iter()
1639 .collect::<DictionaryArray<Int32Type>>();
1640 let run_ends1 = Int32Array::from(vec![1, 2, 3]);
1641 let arr1 = RunArray::try_new(&run_ends1, &dict_values1).unwrap();
1642
1643 let dict_values2 = vec![Some("c"), Some("a")]
1644 .into_iter()
1645 .collect::<DictionaryArray<Int32Type>>();
1646 let run_ends2 = Int32Array::from(vec![1, 2]);
1647 let arr2 = RunArray::try_new(&run_ends2, &dict_values2).unwrap();
1648
1649 let schema = Arc::new(Schema::new(vec![Field::new(
1650 "ree",
1651 arr1.data_type().clone(),
1652 true,
1653 )]));
1654
1655 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1656 let batch2 = RecordBatch::try_new(schema, vec![Arc::new(arr2)]).unwrap();
1657
1658 verify_flight_round_trip(vec![batch1, batch2]).await;
1659 }
1660
1661 #[tokio::test]
1662 async fn test_dictionary_of_struct_of_dict_resend() {
1663 let struct_fields: Vec<Field> = vec![
1667 Field::new_dictionary("dict", DataType::Int32, DataType::Utf8, true),
1668 Field::new("int", DataType::Int32, false),
1669 ];
1670
1671 let inner_values =
1672 StringArray::from(vec![Some("alpha"), None, Some("beta"), Some("gamma")]);
1673 let inner_keys = Int32Array::from_iter_values([0, 1, 2, 3, 0]);
1674 let inner_dict = DictionaryArray::new(inner_keys, Arc::new(inner_values));
1675 let int_array = Int32Array::from(vec![10, 20, 30, 40, 50]);
1676
1677 let struct_array = StructArray::from(vec![
1678 (
1679 Arc::new(struct_fields[0].clone()),
1680 Arc::new(inner_dict) as ArrayRef,
1681 ),
1682 (
1683 Arc::new(struct_fields[1].clone()),
1684 Arc::new(int_array) as ArrayRef,
1685 ),
1686 ]);
1687
1688 let outer_keys = Int8Array::from_iter_values([0, 0, 1, 2]);
1689 let arr1 = DictionaryArray::new(outer_keys, Arc::new(struct_array));
1690
1691 let inner_values2 = StringArray::from(vec![Some("x"), Some("y")]);
1692 let inner_keys2 = Int32Array::from_iter_values([0, 1, 0]);
1693 let inner_dict2 = DictionaryArray::new(inner_keys2, Arc::new(inner_values2));
1694 let int_array2 = Int32Array::from(vec![100, 200, 300]);
1695
1696 let struct_array2 = StructArray::from(vec![
1697 (
1698 Arc::new(struct_fields[0].clone()),
1699 Arc::new(inner_dict2) as ArrayRef,
1700 ),
1701 (
1702 Arc::new(struct_fields[1].clone()),
1703 Arc::new(int_array2) as ArrayRef,
1704 ),
1705 ]);
1706
1707 let outer_keys2 = Int8Array::from_iter_values([0, 1]);
1708 let arr2 = DictionaryArray::new(outer_keys2, Arc::new(struct_array2));
1709
1710 let schema = Arc::new(Schema::new(vec![Field::new(
1711 "dict_struct",
1712 arr1.data_type().clone(),
1713 false,
1714 )]));
1715
1716 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1717 let batch2 = RecordBatch::try_new(schema, vec![Arc::new(arr2)]).unwrap();
1718
1719 verify_flight_round_trip(vec![batch1, batch2]).await;
1720 }
1721
1722 async fn verify_dictionary_list_view_resend<O: OffsetSizeTrait>() {
1723 let mut builder =
1724 GenericListViewBuilder::<O, _>::new(StringDictionaryBuilder::<UInt16Type>::new());
1725
1726 builder.append_value(vec![Some("a"), None, Some("b")]);
1727 let arr1 = builder.finish();
1728
1729 builder.append_value(vec![Some("c"), None, Some("d")]);
1730 let arr2 = builder.finish();
1731
1732 let inner = Arc::new(Field::new_dictionary(
1733 "item",
1734 DataType::UInt16,
1735 DataType::Utf8,
1736 true,
1737 ));
1738 let dt = if O::IS_LARGE {
1739 DataType::LargeListView(inner)
1740 } else {
1741 DataType::ListView(inner)
1742 };
1743 let schema = Arc::new(Schema::new(vec![Field::new("dict_list_view", dt, true)]));
1744
1745 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1746 let batch2 = RecordBatch::try_new(schema, vec![Arc::new(arr2)]).unwrap();
1747
1748 verify_flight_round_trip(vec![batch1, batch2]).await;
1749 }
1750
1751 #[tokio::test]
1752 async fn test_dictionary_list_view_resend() {
1753 verify_dictionary_list_view_resend::<i32>().await;
1754 }
1755
1756 #[tokio::test]
1757 async fn test_dictionary_large_list_view_resend() {
1758 verify_dictionary_list_view_resend::<i64>().await;
1759 }
1760
1761 #[tokio::test]
1762 async fn test_dictionary_fixed_size_list_resend() {
1763 let mut builder =
1764 FixedSizeListBuilder::new(StringDictionaryBuilder::<UInt16Type>::new(), 2);
1765
1766 builder.values().append_value("a");
1767 builder.values().append_value("b");
1768 builder.append(true);
1769 let arr1 = builder.finish();
1770
1771 builder.values().append_value("c");
1772 builder.values().append_value("d");
1773 builder.append(true);
1774 let arr2 = builder.finish();
1775
1776 let schema = Arc::new(Schema::new(vec![Field::new_fixed_size_list(
1777 "dict_fsl",
1778 Field::new_dictionary("item", DataType::UInt16, DataType::Utf8, true),
1779 2,
1780 true,
1781 )]));
1782
1783 let batch1 = RecordBatch::try_new(schema.clone(), vec![Arc::new(arr1)]).unwrap();
1784 let batch2 = RecordBatch::try_new(schema, vec![Arc::new(arr2)]).unwrap();
1785
1786 verify_flight_round_trip(vec![batch1, batch2]).await;
1787 }
1788
1789 async fn verify_flight_round_trip(batches: Vec<RecordBatch>) {
1790 verify_flight_round_trip_with_options(batches, IpcWriteOptions::default()).await;
1791 }
1792
1793 async fn verify_flight_round_trip_with_options(
1796 mut batches: Vec<RecordBatch>,
1797 options: IpcWriteOptions,
1798 ) {
1799 let expected_schema = batches.first().unwrap().schema();
1800
1801 let encoder = FlightDataEncoderBuilder::default()
1802 .with_options(options)
1803 .with_dictionary_handling(DictionaryHandling::Resend)
1804 .build(futures::stream::iter(batches.clone().into_iter().map(Ok)));
1805
1806 let mut expected_batches = batches.drain(..);
1807
1808 let mut decoder = FlightDataDecoder::new(encoder);
1809 while let Some(decoded) = decoder.next().await {
1810 let decoded = decoded.unwrap();
1811 match decoded.payload {
1812 DecodedPayload::None => {}
1813 DecodedPayload::Schema(s) => assert_eq!(s, expected_schema),
1814 DecodedPayload::RecordBatch(b) => {
1815 let expected_batch = expected_batches.next().unwrap();
1816 assert_eq!(b, expected_batch);
1817 }
1818 }
1819 }
1820 }
1821
1822 #[test]
1823 fn test_schema_metadata_encoded() {
1824 let schema = Schema::new(vec![Field::new("data", DataType::Int32, false)]).with_metadata(
1825 HashMap::from([("some_key".to_owned(), "some_value".to_owned())]),
1826 );
1827
1828 let mut dictionary_tracker = DictionaryTracker::new(false);
1829
1830 let got = prepare_schema_for_flight(&schema, &mut dictionary_tracker, false);
1831 assert!(got.metadata().contains_key("some_key"));
1832 }
1833
1834 #[test]
1835 fn test_encode_no_column_batch() {
1836 let batch = RecordBatch::try_new_with_options(
1837 Arc::new(Schema::empty()),
1838 vec![],
1839 &RecordBatchOptions::new().with_row_count(Some(10)),
1840 )
1841 .expect("cannot create record batch");
1842
1843 hydrate_dictionaries(&batch, batch.schema()).expect("failed to optimize");
1844 }
1845
1846 fn make_flight_data(
1847 batch: &RecordBatch,
1848 options: &IpcWriteOptions,
1849 ) -> (Vec<FlightData>, FlightData) {
1850 flight_data_from_arrow_batch(batch, options)
1851 }
1852
1853 fn flight_data_from_arrow_batch(
1854 batch: &RecordBatch,
1855 options: &IpcWriteOptions,
1856 ) -> (Vec<FlightData>, FlightData) {
1857 let data_gen = IpcDataGenerator::default();
1858 let mut dictionary_tracker = DictionaryTracker::new(false);
1859 let mut ipc_write_context = IpcWriteContext::default();
1860
1861 let (encoded_dictionaries, encoded_batch) = data_gen
1862 .encode(
1863 batch,
1864 &mut dictionary_tracker,
1865 options,
1866 &mut ipc_write_context,
1867 )
1868 .expect("DictionaryTracker configured above to not error on replacement");
1869
1870 let flight_dictionaries = encoded_dictionaries.into_iter().map(Into::into).collect();
1871 let flight_batch = encoded_batch.into();
1872
1873 (flight_dictionaries, flight_batch)
1874 }
1875
1876 #[test]
1877 fn test_split_batch_for_grpc_response() {
1878 let max_flight_data_size = 1024;
1879
1880 let c = UInt32Array::from(vec![1, 2, 3, 4, 5, 6]);
1882 let batch = RecordBatch::try_from_iter(vec![("a", Arc::new(c) as ArrayRef)])
1883 .expect("cannot create record batch");
1884 let split: Vec<_> = split_batch_for_grpc_response(batch.clone(), max_flight_data_size);
1885 assert_eq!(split.len(), 1);
1886 assert_eq!(batch, split[0]);
1887
1888 let n_rows = max_flight_data_size + 1;
1890 assert!(n_rows % 2 == 1, "should be an odd number");
1891 let c = UInt8Array::from((0..n_rows).map(|i| (i % 256) as u8).collect::<Vec<_>>());
1892 let batch = RecordBatch::try_from_iter(vec![("a", Arc::new(c) as ArrayRef)])
1893 .expect("cannot create record batch");
1894 let split: Vec<_> = split_batch_for_grpc_response(batch.clone(), max_flight_data_size);
1895 assert_eq!(split.len(), 3);
1896 assert_eq!(
1897 split.iter().map(|batch| batch.num_rows()).sum::<usize>(),
1898 n_rows
1899 );
1900 let a = pretty_format_batches(&split).unwrap().to_string();
1901 let b = pretty_format_batches(&[batch]).unwrap().to_string();
1902 assert_eq!(a, b);
1903 }
1904
1905 #[test]
1906 fn test_split_batch_for_grpc_response_sizes() {
1907 verify_split(2000, 2 * 1024, vec![250, 250, 250, 250, 250, 250, 250, 250]);
1909
1910 verify_split(2000, 4 * 1024, vec![500, 500, 500, 500]);
1912
1913 verify_split(2023, 3 * 1024, vec![337, 337, 337, 337, 337, 337, 1]);
1915
1916 verify_split(10, 1, vec![1, 1, 1, 1, 1, 1, 1, 1, 1, 1]);
1918
1919 verify_split(10, 1024, vec![10]);
1921 }
1922
1923 fn verify_split(
1927 num_input_rows: u64,
1928 max_flight_data_size_bytes: usize,
1929 expected_sizes: Vec<usize>,
1930 ) {
1931 let array: UInt64Array = (0..num_input_rows).collect();
1932
1933 let batch = RecordBatch::try_from_iter(vec![("a", Arc::new(array) as ArrayRef)])
1934 .expect("cannot create record batch");
1935
1936 let input_rows = batch.num_rows();
1937
1938 let split: Vec<_> =
1939 split_batch_for_grpc_response(batch.clone(), max_flight_data_size_bytes);
1940 let sizes: Vec<_> = split.iter().map(RecordBatch::num_rows).collect();
1941 let output_rows: usize = sizes.iter().sum();
1942
1943 assert_eq!(sizes, expected_sizes, "mismatch for {batch:?}");
1944 assert_eq!(input_rows, output_rows, "mismatch for {batch:?}");
1945 }
1946
1947 #[tokio::test]
1951 async fn flight_data_size_even() {
1952 let s1 = StringArray::from_iter_values(std::iter::repeat_n(".10 bytes.", 1024));
1953 let i1 = Int16Array::from_iter_values(0..1024);
1954 let s2 = StringArray::from_iter_values(std::iter::repeat_n("6bytes", 1024));
1955 let i2 = Int64Array::from_iter_values(0..1024);
1956
1957 let batch = RecordBatch::try_from_iter(vec![
1958 ("s1", Arc::new(s1) as _),
1959 ("i1", Arc::new(i1) as _),
1960 ("s2", Arc::new(s2) as _),
1961 ("i2", Arc::new(i2) as _),
1962 ])
1963 .unwrap();
1964
1965 verify_encoded_split(batch, 120).await;
1966 }
1967
1968 #[tokio::test]
1969 async fn flight_data_size_uneven_variable_lengths() {
1970 let array = StringArray::from_iter_values((0..1024).map(|i| "*".repeat(i)));
1972 let batch = RecordBatch::try_from_iter(vec![("data", Arc::new(array) as _)]).unwrap();
1973
1974 verify_encoded_split(batch, 4312).await;
1977 }
1978
1979 #[tokio::test]
1980 async fn flight_data_size_large_row() {
1981 let array1 = StringArray::from_iter_values(vec![
1983 "*".repeat(500),
1984 "*".repeat(500),
1985 "*".repeat(500),
1986 "*".repeat(500),
1987 ]);
1988 let array2 = StringArray::from_iter_values(vec![
1989 "*".to_string(),
1990 "*".repeat(1000),
1991 "*".repeat(2000),
1992 "*".repeat(4000),
1993 ]);
1994
1995 let array3 = StringArray::from_iter_values(vec![
1996 "*".to_string(),
1997 "*".to_string(),
1998 "*".repeat(1000),
1999 "*".repeat(2000),
2000 ]);
2001
2002 let batch = RecordBatch::try_from_iter(vec![
2003 ("a1", Arc::new(array1) as _),
2004 ("a2", Arc::new(array2) as _),
2005 ("a3", Arc::new(array3) as _),
2006 ])
2007 .unwrap();
2008
2009 verify_encoded_split(batch, 5808).await;
2013 }
2014
2015 #[tokio::test]
2016 async fn flight_data_size_string_dictionary() {
2017 let array: DictionaryArray<Int32Type> = (1..1024)
2019 .map(|i| match i % 3 {
2020 0 => Some("value0"),
2021 1 => Some("value1"),
2022 _ => None,
2023 })
2024 .collect();
2025
2026 let batch = RecordBatch::try_from_iter(vec![("a1", Arc::new(array) as _)]).unwrap();
2027
2028 verify_encoded_split(batch, 56).await;
2029 }
2030
2031 #[tokio::test]
2032 async fn flight_data_size_large_dictionary() {
2033 let values: Vec<_> = (1..1024).map(|i| "**".repeat(i)).collect();
2035
2036 let array: DictionaryArray<Int32Type> = values.iter().map(|s| Some(s.as_str())).collect();
2037
2038 let batch = RecordBatch::try_from_iter(vec![("a1", Arc::new(array) as _)]).unwrap();
2039
2040 verify_encoded_split(batch, 3336).await;
2043 }
2044
2045 #[tokio::test]
2046 async fn flight_data_size_large_dictionary_repeated_non_uniform() {
2047 let values = StringArray::from_iter_values((0..1024).map(|i| "******".repeat(i)));
2049 let keys = Int32Array::from_iter_values((0..3000).map(|i| (3000 - i) % 1024));
2050 let array = DictionaryArray::new(keys, Arc::new(values));
2051
2052 let batch = RecordBatch::try_from_iter(vec![("a1", Arc::new(array) as _)]).unwrap();
2053
2054 verify_encoded_split(batch, 5288).await;
2057 }
2058
2059 #[tokio::test]
2060 async fn flight_data_size_multiple_dictionaries() {
2061 let values1: Vec<_> = (1..1024).map(|i| "**".repeat(i)).collect();
2063 let values2: Vec<_> = (1..1024).map(|i| "**".repeat(i % 10)).collect();
2065 let values3: Vec<_> = (1..1024).map(|i| "**".repeat(i % 100)).collect();
2067
2068 let array1: DictionaryArray<Int32Type> = values1.iter().map(|s| Some(s.as_str())).collect();
2069 let array2: DictionaryArray<Int32Type> = values2.iter().map(|s| Some(s.as_str())).collect();
2070 let array3: DictionaryArray<Int32Type> = values3.iter().map(|s| Some(s.as_str())).collect();
2071
2072 let batch = RecordBatch::try_from_iter(vec![
2073 ("a1", Arc::new(array1) as _),
2074 ("a2", Arc::new(array2) as _),
2075 ("a3", Arc::new(array3) as _),
2076 ])
2077 .unwrap();
2078
2079 verify_encoded_split(batch, 4136).await;
2082 }
2083
2084 fn flight_data_size(d: &FlightData) -> usize {
2086 let flight_descriptor_size = d
2087 .flight_descriptor
2088 .as_ref()
2089 .map(|descriptor| {
2090 let path_len: usize = descriptor.path.iter().map(|p| p.len()).sum();
2091
2092 std::mem::size_of_val(descriptor) + descriptor.cmd.len() + path_len
2093 })
2094 .unwrap_or(0);
2095
2096 flight_descriptor_size + d.app_metadata.len() + d.data_body.len() + d.data_header.len()
2097 }
2098
2099 async fn verify_encoded_split(batch: RecordBatch, allowed_overage: usize) {
2115 let num_rows = batch.num_rows();
2116
2117 let mut max_overage_seen = 0;
2119
2120 for max_flight_data_size in [1024, 2021, 5000] {
2121 println!("Encoding {num_rows} with a maximum size of {max_flight_data_size}");
2122
2123 let mut stream = FlightDataEncoderBuilder::new()
2124 .with_max_flight_data_size(max_flight_data_size)
2125 .with_options(IpcWriteOptions::try_new(8, false, MetadataVersion::V5).unwrap())
2127 .build(futures::stream::iter([Ok(batch.clone())]));
2128
2129 let mut i = 0;
2130 while let Some(data) = stream.next().await.transpose().unwrap() {
2131 let actual_data_size = flight_data_size(&data);
2132
2133 let actual_overage = actual_data_size.saturating_sub(max_flight_data_size);
2134
2135 assert!(
2136 actual_overage <= allowed_overage,
2137 "encoded data[{i}]: actual size {actual_data_size}, \
2138 actual_overage: {actual_overage} \
2139 allowed_overage: {allowed_overage}"
2140 );
2141
2142 i += 1;
2143
2144 max_overage_seen = max_overage_seen.max(actual_overage)
2145 }
2146 }
2147
2148 assert_eq!(
2152 allowed_overage, max_overage_seen,
2153 "Specified overage was too high"
2154 );
2155 }
2156}