Skip to main content

arrow_flight/
encode.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use 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/// Creates a [`Stream`] of [`FlightData`]s from a
30/// `Stream` of [`Result`]<[`RecordBatch`], [`FlightError`]>.
31///
32/// This can be used to implement [`FlightService::do_get`] in an
33/// Arrow Flight implementation;
34///
35/// This structure encodes a stream of `Result`s rather than `RecordBatch`es  to
36/// propagate errors from streaming execution, where the generation of the
37/// `RecordBatch`es is incremental, and an error may occur even after
38/// several have already been successfully produced.
39///
40/// # Caveats
41/// 1. When [`DictionaryHandling`] is [`DictionaryHandling::Hydrate`],
42///    [`DictionaryArray`]s are converted to their underlying types prior to
43///    transport.
44///    When [`DictionaryHandling`] is [`DictionaryHandling::Resend`], Dictionary [`FlightData`] is sent with every
45///    [`RecordBatch`] that contains a [`DictionaryArray`](arrow_array::array::DictionaryArray).
46///    See <https://github.com/apache/arrow-rs/issues/3389>.
47///
48/// [`DictionaryArray`]: arrow_array::array::DictionaryArray
49///
50/// # Example
51/// ```no_run
52/// # use std::sync::Arc;
53/// # use arrow_array::{ArrayRef, RecordBatch, UInt32Array};
54/// # async fn f() {
55/// # let c1 = UInt32Array::from(vec![1, 2, 3, 4, 5, 6]);
56/// # let batch = RecordBatch::try_from_iter(vec![
57/// #      ("a", Arc::new(c1) as ArrayRef)
58/// #   ])
59/// #   .expect("cannot create record batch");
60/// use arrow_flight::encode::FlightDataEncoderBuilder;
61///
62/// // Get an input stream of Result<RecordBatch, FlightError>
63/// let input_stream = futures::stream::iter(vec![Ok(batch)]);
64///
65/// // Build a stream of `Result<FlightData>` (e.g. to return for do_get)
66/// let flight_data_stream = FlightDataEncoderBuilder::new()
67///  .build(input_stream);
68///
69/// // Create a tonic `Response` that can be returned from a Flight server
70/// let response = tonic::Response::new(flight_data_stream);
71/// # }
72/// ```
73///
74/// # Example: Sending `Vec<RecordBatch>`
75///
76/// You can create a [`Stream`] to pass to [`Self::build`] from an existing
77/// `Vec` of `RecordBatch`es like this:
78///
79/// ```
80/// # use std::sync::Arc;
81/// # use arrow_array::{ArrayRef, RecordBatch, UInt32Array};
82/// # async fn f() {
83/// # fn make_batches() -> Vec<RecordBatch> {
84/// #   let c1 = UInt32Array::from(vec![1, 2, 3, 4, 5, 6]);
85/// #   let batch = RecordBatch::try_from_iter(vec![
86/// #      ("a", Arc::new(c1) as ArrayRef)
87/// #   ])
88/// #   .expect("cannot create record batch");
89/// #   vec![batch.clone(), batch.clone()]
90/// # }
91/// use arrow_flight::encode::FlightDataEncoderBuilder;
92///
93/// // Get batches that you want to send via Flight
94/// let batches: Vec<RecordBatch> = make_batches();
95///
96/// // Create an input stream of Result<RecordBatch, FlightError>
97/// let input_stream = futures::stream::iter(
98///   batches.into_iter().map(Ok)
99/// );
100///
101/// // Build a stream of `Result<FlightData>` (e.g. to return for do_get)
102/// let flight_data_stream = FlightDataEncoderBuilder::new()
103///  .build(input_stream);
104/// # }
105/// ```
106///
107/// # Example: Determining schema of encoded data
108///
109/// Encoding flight data may hydrate dictionaries, see [`DictionaryHandling`] for more information,
110/// which changes the schema of the encoded data compared to the input record batches.
111/// The fully hydrated schema can be accessed using the [`FlightDataEncoder::known_schema`] method
112/// and explicitly informing the builder of the schema using [`FlightDataEncoderBuilder::with_schema`].
113///
114/// ```
115/// # use std::sync::Arc;
116/// # use arrow_array::{ArrayRef, RecordBatch, UInt32Array};
117/// # async fn f() {
118/// # let c1 = UInt32Array::from(vec![1, 2, 3, 4, 5, 6]);
119/// # let batch = RecordBatch::try_from_iter(vec![
120/// #      ("a", Arc::new(c1) as ArrayRef)
121/// #   ])
122/// #   .expect("cannot create record batch");
123/// use arrow_flight::encode::FlightDataEncoderBuilder;
124///
125/// // Get the schema of the input stream
126/// let schema = batch.schema();
127///
128/// // Get an input stream of Result<RecordBatch, FlightError>
129/// let input_stream = futures::stream::iter(vec![Ok(batch)]);
130///
131/// // Build a stream of `Result<FlightData>` (e.g. to return for do_get)
132/// let flight_data_stream = FlightDataEncoderBuilder::new()
133///  // Inform the builder of the input stream schema
134///  .with_schema(schema)
135///  .build(input_stream);
136///
137/// // Retrieve the schema of the encoded data
138/// let encoded_schema = flight_data_stream.known_schema();
139/// # }
140/// ```
141///
142/// [`FlightService::do_get`]: crate::flight_service_server::FlightService::do_get
143/// [`FlightError`]: crate::error::FlightError
144#[derive(Debug)]
145pub struct FlightDataEncoderBuilder {
146    /// The maximum approximate target message size in bytes
147    /// (see details on [`Self::with_max_flight_data_size`]).
148    max_flight_data_size: usize,
149    /// Ipc writer options
150    options: IpcWriteOptions,
151    /// Metadata to add to the schema message
152    app_metadata: Bytes,
153    /// Optional schema, if known before data.
154    schema: Option<SchemaRef>,
155    /// Optional flight descriptor, if known before data.
156    descriptor: Option<FlightDescriptor>,
157    /// Deterimines how `DictionaryArray`s are encoded for transport.
158    /// See [`DictionaryHandling`] for more information.
159    dictionary_handling: DictionaryHandling,
160}
161
162/// Default target size for encoded [`FlightData`].
163///
164/// Note this value would normally be 4MB, but the size calculation is
165/// somewhat inexact, so we set it to 2MB.
166pub 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    /// Create a new [`FlightDataEncoderBuilder`].
183    pub fn new() -> Self {
184        Self::default()
185    }
186
187    /// Set the (approximate) maximum size, in bytes, of the
188    /// [`FlightData`] produced by this encoder. Defaults to 2MB.
189    ///
190    /// Since there is often a maximum message size for gRPC messages
191    /// (typically around 4MB), this encoder splits up [`RecordBatch`]s
192    /// (preserving order) into multiple [`FlightData`] objects to
193    /// limit the size individual messages sent via gRPC.
194    ///
195    /// The size is approximate because of the additional encoding
196    /// overhead on top of the underlying data buffers themselves.
197    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    /// Set [`DictionaryHandling`] for encoder
203    pub fn with_dictionary_handling(mut self, dictionary_handling: DictionaryHandling) -> Self {
204        self.dictionary_handling = dictionary_handling;
205        self
206    }
207
208    /// Specify application specific metadata included in the
209    /// [`FlightData::app_metadata`] field of the the first Schema
210    /// message
211    pub fn with_metadata(mut self, app_metadata: Bytes) -> Self {
212        self.app_metadata = app_metadata;
213        self
214    }
215
216    /// Set the [`IpcWriteOptions`] used to encode the [`RecordBatch`]es for transport.
217    pub fn with_options(mut self, options: IpcWriteOptions) -> Self {
218        self.options = options;
219        self
220    }
221
222    /// Specify a schema for the RecordBatches being sent. If a schema
223    /// is not specified, an encoded Schema message will be sent when
224    /// the first [`RecordBatch`], if any, is encoded. Some clients
225    /// expect a Schema message even if there is no data sent.
226    pub fn with_schema(mut self, schema: SchemaRef) -> Self {
227        self.schema = Some(schema);
228        self
229    }
230
231    /// Specify a flight descriptor in the first FlightData message.
232    pub fn with_flight_descriptor(mut self, descriptor: Option<FlightDescriptor>) -> Self {
233        self.descriptor = descriptor;
234        self
235    }
236
237    /// Takes a [`Stream`] of [`Result<RecordBatch>`] and returns a [`Stream`]
238    /// of [`FlightData`], consuming self.
239    ///
240    /// See example on [`Self`] and [`FlightDataEncoder`] for more details
241    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
266/// Stream that encodes a stream of record batches to flight data.
267///
268/// See [`FlightDataEncoderBuilder`] for details and example.
269pub struct FlightDataEncoder {
270    /// Input stream
271    inner: BoxStream<'static, Result<RecordBatch>>,
272    /// schema, set after the first batch
273    schema: Option<SchemaRef>,
274    /// Target maximum size of flight data
275    /// (see details on [`FlightDataEncoderBuilder::with_max_flight_data_size`]).
276    max_flight_data_size: usize,
277    /// do the encoding / tracking of dictionaries
278    encoder: FlightIpcEncoder,
279    /// optional metadata to add to schema FlightData
280    app_metadata: Option<Bytes>,
281    /// data queued up to send but not yet sent
282    queue: VecDeque<FlightData>,
283    /// Is this stream done (inner is empty or errored)
284    done: bool,
285    /// cleared after the first FlightData message is sent
286    descriptor: Option<FlightDescriptor>,
287    /// Deterimines how `DictionaryArray`s are encoded for transport.
288    /// See [`DictionaryHandling`] for more information.
289    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 schema is known up front, enqueue it immediately
318        if let Some(schema) = schema {
319            encoder.encode_schema(&schema);
320        }
321
322        encoder
323    }
324
325    /// Report the schema of the encoded data when known.
326    /// A schema is known when provided via the [`FlightDataEncoderBuilder::with_schema`] method.
327    pub fn known_schema(&self) -> Option<SchemaRef> {
328        self.schema.clone()
329    }
330
331    /// Place the `FlightData` in the queue to send
332    #[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    /// Encodes schema as a [`FlightData`] in self.queue.
341    /// Updates `self.schema` and returns the new schema
342    fn encode_schema(&mut self, schema: &SchemaRef) -> SchemaRef {
343        // The first message is the schema message, and all
344        // batches have the same schema
345        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        // attach any metadata requested
354        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        // remember schema
359        self.schema = Some(schema.clone());
360        schema
361    }
362
363    /// Encodes batch into one or more `FlightData` messages in self.queue
364    fn encode_batch(&mut self, batch: RecordBatch) -> Result<()> {
365        let schema = match &self.schema {
366            Some(schema) => schema.clone(),
367            // encode the schema if this is the first time we have seen it
368            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); // handle empty batches  
378        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            // Any messages queued to send?
406            if let Some(data) = self.queue.pop_front() {
407                return Poll::Ready(Some(Ok(data)));
408            }
409
410            // Get next batch
411            let batch = ready!(self.inner.poll_next_unpin(cx));
412
413            match batch {
414                None => {
415                    // inner is done
416                    self.done = true;
417                    // queue must also be empty so we are done
418                    assert!(self.queue.is_empty());
419                    return Poll::Ready(None);
420                }
421                Some(Err(e)) => {
422                    // error from inner
423                    self.done = true;
424                    self.queue.clear();
425                    return Poll::Ready(Some(Err(e)));
426                }
427                Some(Ok(batch)) => {
428                    // had data, encode into the queue
429                    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/// Defines how a [`FlightDataEncoder`] encodes [`DictionaryArray`]s
441///
442/// [`DictionaryArray`]: arrow_array::DictionaryArray
443///
444/// In the arrow flight protocol dictionary values and keys are sent as two separate messages.
445/// When a sender is encoding a [`RecordBatch`] containing ['DictionaryArray'] columns, it will
446/// first send a dictionary batch (a batch with header `MessageHeader::DictionaryBatch`) containing
447/// the dictionary values. The receiver is responsible for reading this batch and maintaining state that associates
448/// those dictionary values with the corresponding array using the `dict_id` as a key.
449///
450/// After sending the dictionary batch the sender will send the array data in a batch with header `MessageHeader::RecordBatch`.
451/// For any dictionary array batches in this message, the encoded flight message will only contain the dictionary keys. The receiver
452/// is then responsible for rebuilding the `DictionaryArray` on the client side using the dictionary values from the DictionaryBatch message
453/// and the keys from the RecordBatch message.
454///
455/// For example, if we have a batch with a `TypedDictionaryArray<'_, UInt32Type, Utf8Type>` (a dictionary array where they keys are `u32` and the
456/// values are `String`), then the DictionaryBatch will contain a `StringArray` and the RecordBatch will contain a `UInt32Array`.
457///
458/// Note that since `dict_id` defined in the `Schema` is used as a key to associate dictionary values to their arrays it is required that each
459/// `DictionaryArray` in a `RecordBatch` have a unique `dict_id`.
460///
461/// The current implementation does not support "delta" dictionaries so a new dictionary batch will be sent each time the encoder sees a
462/// dictionary which is not pointer-equal to the previously observed dictionary for a given `dict_id`.
463///
464/// For clients which may not support `DictionaryEncoding`, the `DictionaryHandling::Hydrate` method will bypass the process defined above
465/// and "hydrate" any `DictionaryArray` in the batch to their underlying value type (e.g. `TypedDictionaryArray<'_, UInt32Type, Utf8Type>` will
466/// be sent as a `StringArray`). With this method all data will be sent in ``MessageHeader::RecordBatch` messages and the batch schema
467/// will be adjusted so that all dictionary encoded fields are changed to fields of the dictionary value type.
468#[derive(Debug, PartialEq)]
469pub enum DictionaryHandling {
470    /// Expands to the underlying type (default). This likely sends more data
471    /// over the network but requires less memory (dictionaries are not tracked)
472    /// and is more compatible with other arrow flight client implementations
473    /// that may not support `DictionaryEncoding`
474    ///
475    /// See also:
476    /// * <https://github.com/apache/arrow-rs/issues/1206>
477    Hydrate,
478    /// Send dictionary FlightData with every RecordBatch that contains a
479    /// [`DictionaryArray`]. See [`Self::Hydrate`] for more tradeoffs. No
480    /// attempt is made to skip sending the same (logical) dictionary values
481    /// twice.
482    ///
483    /// [`DictionaryArray`]: arrow_array::DictionaryArray
484    ///
485    /// This requires identifying the different dictionaries in use and assigning
486    //  them unique IDs
487    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                // Recurse into value type to handle nested dicts being stripped
532                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                // Recurse into value type BEFORE registering this dict's id,
545                // matching the depth-first order of encode_dictionaries in the
546                // IPC writer which processes nested dicts before the parent.
547                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
646/// Prepare an arrow Schema for transport over the Arrow Flight protocol
647///
648/// Convert dictionary types to underlying types
649///
650/// See hydrate_dictionary for more information
651fn 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
665/// Split [`RecordBatch`] so it hopefully fits into a gRPC response.
666///
667/// Data is zero-copy sliced into batches.
668///
669/// Note: this method does not take into account already sliced
670/// arrays: <https://github.com/apache/arrow-rs/issues/3407>
671fn 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
697/// The data needed to encode a stream of flight data, holding on to
698/// shared Dictionaries.
699///
700/// TODO: at allow dictionaries to be flushed / avoid building them
701///
702/// TODO limit on the number of dictionaries???
703struct 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    /// Encode a schema as a FlightData
721    fn encode_schema(&self, schema: &Schema) -> FlightData {
722        SchemaAsIpc::new(schema, &self.options).into()
723    }
724
725    /// Convert a `RecordBatch` to a Vec of `FlightData` representing
726    /// dictionaries and a `FlightData` representing the batch
727    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
745/// Hydrates any dictionaries arrays in `batch` to its underlying type. See
746/// hydrate_dictionary for more information.
747fn 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
762/// Hydrates a dictionary to its underlying type.
763fn 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    /// ensure only the batch's used data (not the allocated data) is sent
808    /// <https://github.com/apache/arrow-rs/issues/208>
809    fn test_encode_flight_data() {
810        // use 8-byte alignment - default alignment is 64 which produces bigger ipc data
811        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        // Round trip a batch through Flight with IPC body compression enabled. This exercises
902        // the compressed `IpcDataGenerator::encode` path (per-buffer codec output), which the
903        // uncompressed Flight tests and the writer-based compression tests do not cover.
904        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        // Create a schema with two dictionary fields that have the same dict ID
970        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        // {"k1":"a","k2":null,"k3":"b"}
1483        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        // {"k1":"c","k2":null,"k3":"d"}
1494        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        // array without dictionary fields
1543        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        // {"k1":"a","k2":null,"k3":"b"}
1589        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        // {"k1":"c","k2":null,"k3":"d"}
1600        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        // Dict(Int8, Struct { dict: Dict(Int32, Utf8), int: Int32 })
1664        // This exercises the Dictionary branch recursing into its value type
1665        // before assigning its own dict_id (depth-first ordering).
1666        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    /// Encode `batches` through a [`FlightDataEncoderBuilder`] using `options`, decode them
1794    /// again, and assert the decoded batches match the originals.
1795    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        // no split
1881        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        // split once
1889        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        // 2000 8 byte entries into 2k pieces: 8 chunks of 250 rows
1908        verify_split(2000, 2 * 1024, vec![250, 250, 250, 250, 250, 250, 250, 250]);
1909
1910        // 2000 8 byte entries into 4k pieces: 4 chunks of 500 rows
1911        verify_split(2000, 4 * 1024, vec![500, 500, 500, 500]);
1912
1913        // 2023 8 byte entries into 3k pieces does not divide evenly
1914        verify_split(2023, 3 * 1024, vec![337, 337, 337, 337, 337, 337, 1]);
1915
1916        // 10 8 byte entries into 1 byte pieces means each rows gets its own
1917        verify_split(10, 1, vec![1, 1, 1, 1, 1, 1, 1, 1, 1, 1]);
1918
1919        // 10 8 byte entries into 1k byte pieces means one piece
1920        verify_split(10, 1024, vec![10]);
1921    }
1922
1923    /// Creates a UInt64Array of 8 byte integers with input_rows rows
1924    /// `max_flight_data_size_bytes` pieces and verifies the row counts in
1925    /// those pieces
1926    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    // test sending record batches
1948    // test sending record batches with multiple different dictionaries
1949
1950    #[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        // each row has a longer string than the last with increasing lengths 0 --> 1024
1971        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        // overage is much higher than ideal
1975        // https://github.com/apache/arrow-rs/issues/3478
1976        verify_encoded_split(batch, 4312).await;
1977    }
1978
1979    #[tokio::test]
1980    async fn flight_data_size_large_row() {
1981        // batch with individual that can each exceed the batch size
1982        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        // 5k over limit (which is 2x larger than limit of 5k)
2010        // overage is much higher than ideal
2011        // https://github.com/apache/arrow-rs/issues/3478
2012        verify_encoded_split(batch, 5808).await;
2013    }
2014
2015    #[tokio::test]
2016    async fn flight_data_size_string_dictionary() {
2017        // Small dictionary (only 2 distinct values ==> 2 entries in dictionary)
2018        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        // large dictionary (all distinct values ==> 1024 entries in dictionary)
2034        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        // overage is much higher than ideal
2041        // https://github.com/apache/arrow-rs/issues/3478
2042        verify_encoded_split(batch, 3336).await;
2043    }
2044
2045    #[tokio::test]
2046    async fn flight_data_size_large_dictionary_repeated_non_uniform() {
2047        // large dictionary (1024 distinct values) that are used throughout the array
2048        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        // overage is much higher than ideal
2055        // https://github.com/apache/arrow-rs/issues/3478
2056        verify_encoded_split(batch, 5288).await;
2057    }
2058
2059    #[tokio::test]
2060    async fn flight_data_size_multiple_dictionaries() {
2061        // high cardinality
2062        let values1: Vec<_> = (1..1024).map(|i| "**".repeat(i)).collect();
2063        // highish cardinality
2064        let values2: Vec<_> = (1..1024).map(|i| "**".repeat(i % 10)).collect();
2065        // medium cardinality
2066        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        // overage is much higher than ideal
2080        // https://github.com/apache/arrow-rs/issues/3478
2081        verify_encoded_split(batch, 4136).await;
2082    }
2083
2084    /// Return size, in memory of flight data
2085    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    /// Coverage for <https://github.com/apache/arrow-rs/issues/3478>
2100    ///
2101    /// Encodes the specified batch using several values of
2102    /// `max_flight_data_size` between 1K to 5K and ensures that the
2103    /// resulting size of the flight data stays within the limit
2104    /// + `allowed_overage`
2105    ///
2106    /// `allowed_overage` is how far off the actual data encoding is
2107    /// from the target limit that was set. It is an improvement when
2108    /// the allowed_overage decreses.
2109    ///
2110    /// Note this overhead will likely always be greater than zero to
2111    /// account for encoding overhead such as IPC headers and padding.
2112    ///
2113    ///
2114    async fn verify_encoded_split(batch: RecordBatch, allowed_overage: usize) {
2115        let num_rows = batch.num_rows();
2116
2117        // Track the overall required maximum overage
2118        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                // use 8-byte alignment - default alignment is 64 which produces bigger ipc data
2126                .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        // ensure that the specified overage is exactly the maxmium so
2149        // that when the splitting logic improves, the tests must be
2150        // updated to reflect the better logic
2151        assert_eq!(
2152            allowed_overage, max_overage_seen,
2153            "Specified overage was too high"
2154        );
2155    }
2156}