Skip to main content

client/
flight.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use 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 Flight peer associated with this response stream.
25    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}