1use arrow_flight::FlightData;
16use common_grpc::flight::{FlightDecoder, FlightMessage};
17use futures_util::{Stream, StreamExt};
18use snafu::{OptionExt, ResultExt};
19
20use crate::Result;
21use crate::error::{ConvertFlightDataSnafu, Error, IllegalFlightMessagesSnafu};
22
23pub(crate) struct FlightMessageReader<S: Stream + Unpin> {
24 remote_addr: String,
26 messages: S,
27}
28
29impl<S> FlightMessageReader<S>
30where
31 S: Stream<Item = Result<FlightMessage>> + Unpin,
32{
33 pub(crate) fn new(remote_addr: impl Into<String>, messages: S) -> Self {
34 Self {
35 remote_addr: remote_addr.into(),
36 messages,
37 }
38 }
39
40 pub(crate) fn remote_addr(&self) -> &str {
41 &self.remote_addr
42 }
43
44 pub(crate) async fn read_first(&mut self) -> Result<FlightMessage> {
45 self.read_next().await?.context(IllegalFlightMessagesSnafu {
46 reason: "Expect the response not to be empty",
47 })
48 }
49
50 pub(crate) async fn read_next(&mut self) -> Result<Option<FlightMessage>> {
51 self.messages.next().await.transpose()
52 }
53}
54
55pub(crate) fn decode_flight_data(
56 decoder: &mut FlightDecoder,
57 flight_data: std::result::Result<FlightData, tonic::Status>,
58) -> Option<Result<FlightMessage>> {
59 flight_data
60 .map_err(Error::from)
61 .and_then(|data| decoder.try_decode(&data).context(ConvertFlightDataSnafu))
62 .transpose()
63}
64
65#[cfg(test)]
66mod tests {
67 use std::sync::Arc;
68
69 use common_grpc::flight::FlightEncoder;
70 use datatypes::arrow::array::{DictionaryArray, StringArray, UInt32Array};
71 use datatypes::arrow::datatypes::{DataType, Field, Schema, UInt32Type};
72 use datatypes::arrow::record_batch::RecordBatch;
73
74 use super::*;
75
76 #[test]
77 fn test_decode_flight_data_skips_dictionary_batches() {
78 let schema = Arc::new(Schema::new(vec![Field::new_dictionary(
79 "host",
80 DataType::UInt32,
81 DataType::Utf8,
82 false,
83 )]));
84 let batch = RecordBatch::try_new(
85 schema.clone(),
86 vec![Arc::new(DictionaryArray::<UInt32Type>::new(
87 UInt32Array::from(vec![0, 1, 0]),
88 Arc::new(StringArray::from(vec!["host-a", "host-b"])),
89 ))],
90 )
91 .unwrap();
92
93 let mut encoder = FlightEncoder::default();
94 let mut flight_data = Vec::new();
95 flight_data.extend(encoder.encode(FlightMessage::Schema(schema.clone())));
96 let encoded_batch = encoder.encode(FlightMessage::RecordBatch(batch.clone()));
97 assert_eq!(2, encoded_batch.len());
98 flight_data.extend(encoded_batch);
99
100 let mut decoder = FlightDecoder::default();
101 let messages = flight_data
102 .into_iter()
103 .filter_map(|data| decode_flight_data(&mut decoder, Ok(data)))
104 .collect::<Result<Vec<_>>>()
105 .unwrap();
106
107 assert_eq!(2, messages.len());
108 assert!(matches!(&messages[0], FlightMessage::Schema(actual) if actual == &schema));
109 assert!(matches!(&messages[1], FlightMessage::RecordBatch(actual) if actual == &batch));
110 }
111}