Skip to main content

client/
region.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 std::sync::Arc;
16use std::time::Duration;
17
18use api::region::RegionResponse;
19use api::v1::ResponseHeader;
20use api::v1::region::{
21    RegionRequest, RegionRequestHeader, RemoteDynFilterRequest, RemoteDynFilterUnregister,
22    RemoteDynFilterUpdate, region_request, remote_dyn_filter_request,
23};
24use arc_swap::ArcSwapOption;
25use arrow_flight::Ticket;
26use async_stream::stream;
27use async_trait::async_trait;
28use common_error::ext::{BoxedError, ErrorExt};
29use common_error::status_code::StatusCode;
30use common_grpc::flight::{FlightDecoder, FlightMessage};
31use common_meta::error::{self as meta_error, Result as MetaResult};
32use common_meta::node_manager::Datanode;
33use common_query::request::QueryRequest;
34use common_recordbatch::error::ExternalSnafu;
35use common_recordbatch::{RecordBatch, RecordBatchStreamWrapper, SendableRecordBatchStream};
36use common_telemetry::error;
37use common_telemetry::tracing::Span;
38use common_telemetry::tracing_context::TracingContext;
39use futures_util::Stream;
40use prost::Message;
41use query::query_engine::DefaultSerializer;
42use snafu::{OptionExt, ResultExt, location};
43use substrait::{DFLogicalSubstraitConvertor, SubstraitPlan};
44use tokio_stream::StreamExt;
45use tonic::codec::CompressionEncoding;
46
47use crate::error::{
48    self, FlightGetSnafu, IllegalDatabaseResponseSnafu, IllegalFlightMessagesSnafu,
49    MissingFieldSnafu, Result, ServerSnafu,
50};
51use crate::flight::{FlightMessageReader, decode_flight_data};
52use crate::{Client, metrics};
53
54const FLIGHT_DO_GET_TIMEOUT: Duration = Duration::from_secs(10);
55
56#[derive(Debug)]
57pub struct RegionRequester {
58    client: Client,
59    send_compression: bool,
60    accept_compression: bool,
61}
62
63#[async_trait]
64impl Datanode for RegionRequester {
65    async fn handle(&self, request: RegionRequest) -> MetaResult<RegionResponse> {
66        self.handle_inner(request).await.map_err(|err| {
67            if err.should_retry() {
68                meta_error::Error::RetryLater {
69                    source: BoxedError::new(err),
70                    clean_poisons: false,
71                }
72            } else {
73                meta_error::Error::External {
74                    source: BoxedError::new(err),
75                    location: location!(),
76                }
77            }
78        })
79    }
80
81    async fn handle_query(&self, request: QueryRequest) -> MetaResult<SendableRecordBatchStream> {
82        let plan = DFLogicalSubstraitConvertor
83            .encode(&request.plan, DefaultSerializer)
84            .map_err(BoxedError::new)
85            .context(meta_error::ExternalSnafu)?
86            .to_vec();
87        let request = api::v1::region::QueryRequest {
88            header: request.header,
89            region_id: request.region_id.as_u64(),
90            plan,
91        };
92
93        let ticket = Ticket {
94            ticket: request.encode_to_vec().into(),
95        };
96        self.do_get_inner(ticket)
97            .await
98            .map_err(BoxedError::new)
99            .context(meta_error::ExternalSnafu)
100    }
101}
102
103impl RegionRequester {
104    pub fn new(client: Client, send_compression: bool, accept_compression: bool) -> Self {
105        Self {
106            client,
107            send_compression,
108            accept_compression,
109        }
110    }
111
112    pub async fn do_get_inner(&self, ticket: Ticket) -> Result<SendableRecordBatchStream> {
113        let mut flight_client = self
114            .client
115            .make_flight_client(self.send_compression, self.accept_compression)?;
116        // Limit Flight DoGet response time without limiting query stream execution.
117        let addr = flight_client.addr().to_string();
118        let mut request = tonic::Request::new(ticket);
119        request.set_timeout(FLIGHT_DO_GET_TIMEOUT);
120        let response = flight_client
121            .mut_inner()
122            .do_get(request)
123            .await
124            .or_else(|e| {
125                let tonic_code = e.code();
126                let e: error::Error = e.into();
127                error!(
128                    e; "Failed to do Flight get, addr: {}, code: {}",
129                    addr,
130                    tonic_code
131                );
132                Err(BoxedError::new(e)).with_context(|_| FlightGetSnafu {
133                    addr: addr.clone(),
134                    tonic_code,
135                })
136            })?;
137
138        let flight_data_stream = response.into_inner();
139        let mut decoder = FlightDecoder::default();
140
141        let flight_message_stream = flight_data_stream
142            .filter_map(move |flight_data| decode_flight_data(&mut decoder, flight_data));
143
144        recordbatches_from_flight_message_stream(addr, flight_message_stream).await
145    }
146
147    async fn handle_inner(&self, request: RegionRequest) -> Result<RegionResponse> {
148        let request_body = request
149            .body
150            .as_ref()
151            .with_context(|| MissingFieldSnafu { field: "body" })?;
152        let is_insert = matches!(
153            request_body,
154            region_request::Body::Inserts(_) | region_request::Body::BulkInsert(_)
155        );
156        let request_type = request_body.as_ref().to_string();
157        let _timer = metrics::METRIC_REGION_REQUEST_GRPC
158            .with_label_values(&[request_type.as_str()])
159            .start_timer();
160
161        let (addr, mut client) = self.client.raw_region_client()?;
162        if is_insert {
163            if self.send_compression {
164                client = client.send_compressed(CompressionEncoding::Zstd);
165            }
166            if self.accept_compression {
167                client = client.accept_compressed(CompressionEncoding::Zstd);
168            }
169        }
170
171        let response = client
172            .handle(request)
173            .await
174            .map_err(|e| {
175                let code = e.code();
176                // Uses `Error::RegionServer` instead of `Error::Server`
177                error::Error::RegionServer {
178                    addr,
179                    code,
180                    source: BoxedError::new(error::Error::from(e)),
181                    location: location!(),
182                }
183            })?
184            .into_inner();
185
186        check_response_header(&response.header)?;
187
188        Ok(RegionResponse::from_region_response(response))
189    }
190
191    pub async fn handle(&self, request: RegionRequest) -> Result<RegionResponse> {
192        self.handle_inner(request).await
193    }
194
195    pub async fn handle_remote_dyn_filter_update(
196        &self,
197        query_id: impl Into<String>,
198        update: RemoteDynFilterUpdate,
199    ) -> Result<RegionResponse> {
200        self.handle_inner(build_remote_dyn_filter_update_request(query_id, update))
201            .await
202    }
203
204    pub async fn handle_remote_dyn_filter_unregister(
205        &self,
206        query_id: impl Into<String>,
207        unregister: RemoteDynFilterUnregister,
208    ) -> Result<RegionResponse> {
209        self.handle_inner(build_remote_dyn_filter_unregister_request(
210            query_id, unregister,
211        ))
212        .await
213    }
214}
215
216async fn recordbatches_from_flight_message_stream<S>(
217    addr: String,
218    flight_message_stream: S,
219) -> Result<SendableRecordBatchStream>
220where
221    S: Stream<Item = Result<FlightMessage>> + Send + Unpin + 'static,
222{
223    let mut reader = FlightMessageReader::new(addr.clone(), flight_message_stream);
224    let FlightMessage::Schema(schema) = reader
225        .read_first()
226        .await
227        .map_err(|error| flight_stream_error(reader.remote_addr(), error))?
228    else {
229        return IllegalFlightMessagesSnafu {
230            reason: "Expect schema to be the first flight message",
231        }
232        .fail()
233        .map_err(|error| flight_stream_error(reader.remote_addr(), error));
234    };
235
236    let metrics = Arc::new(ArcSwapOption::from(None));
237    let metrics_ref = metrics.clone();
238
239    let tracing_context = TracingContext::from_current_span();
240
241    let schema =
242        Arc::new(datatypes::schema::Schema::try_from(schema).context(error::ConvertSchemaSnafu)?);
243    let schema_cloned = schema.clone();
244    let stream_addr = addr;
245    let stream = Box::pin(stream!({
246        let _span = tracing_context.attach(common_telemetry::tracing::info_span!(
247            "poll_flight_data_stream"
248        ));
249
250        loop {
251            let flight_message = match reader.read_next().await {
252                Ok(Some(message)) => message,
253                Ok(None) => break,
254                Err(error) => {
255                    yield Err(BoxedError::new(flight_stream_error(&stream_addr, error)))
256                        .context(ExternalSnafu);
257                    break;
258                }
259            };
260
261            match flight_message {
262                FlightMessage::RecordBatch(record_batch) => {
263                    // Deliver each batch immediately. In particular, do not
264                    // wait for a possible following Metrics message; it is
265                    // consumed on the next poll of this stream.
266                    yield Ok(RecordBatch::from_df_record_batch(
267                        schema_cloned.clone(),
268                        record_batch,
269                    ));
270                }
271                FlightMessage::Metrics(s) => {
272                    // Metrics may arrive before the next RecordBatch.
273                    match serde_json::from_str(&s) {
274                        Ok(metrics) => {
275                            metrics_ref.swap(Some(Arc::new(metrics)));
276                        }
277                        Err(error) => {
278                            common_telemetry::warn!(
279                                "Failed to decode region Flight metrics: {}",
280                                error
281                            );
282                        }
283                    }
284                    continue;
285                }
286                _ => {
287                    yield IllegalFlightMessagesSnafu {
288                        reason: "A Schema message must be succeeded exclusively by a set of RecordBatch messages"
289                    }
290                    .fail()
291                    .map_err(BoxedError::new)
292                    .context(ExternalSnafu);
293                    break;
294                }
295            }
296        }
297    }));
298    let record_batch_stream = RecordBatchStreamWrapper {
299        schema,
300        stream,
301        output_ordering: None,
302        metrics,
303        span: Span::current(),
304    };
305    Ok(Box::pin(record_batch_stream))
306}
307
308fn flight_stream_error(addr: &str, error: error::Error) -> error::Error {
309    let tonic_code = error.tonic_code().unwrap_or(tonic::Code::Unknown);
310    if error.status_code().should_log_error() {
311        error!(
312            error; "Failed to receive Flight data, addr: {}, code: {}",
313            addr,
314            tonic_code
315        );
316    }
317
318    error::Error::FlightGet {
319        addr: addr.to_string(),
320        tonic_code,
321        source: BoxedError::new(error),
322    }
323}
324
325pub fn build_remote_dyn_filter_update_request(
326    query_id: impl Into<String>,
327    update: RemoteDynFilterUpdate,
328) -> RegionRequest {
329    build_remote_dyn_filter_request(
330        query_id.into(),
331        remote_dyn_filter_request::Action::Update(update),
332    )
333}
334
335pub fn build_remote_dyn_filter_unregister_request(
336    query_id: impl Into<String>,
337    unregister: RemoteDynFilterUnregister,
338) -> RegionRequest {
339    build_remote_dyn_filter_request(
340        query_id.into(),
341        remote_dyn_filter_request::Action::Unregister(unregister),
342    )
343}
344
345fn build_remote_dyn_filter_request(
346    query_id: String,
347    action: remote_dyn_filter_request::Action,
348) -> RegionRequest {
349    RegionRequest {
350        header: Some(RegionRequestHeader {
351            tracing_context: TracingContext::from_current_span().to_w3c(),
352            ..Default::default()
353        }),
354        body: Some(region_request::Body::RemoteDynFilter(
355            RemoteDynFilterRequest {
356                query_id,
357                action: Some(action),
358            },
359        )),
360    }
361}
362
363pub fn check_response_header(header: &Option<ResponseHeader>) -> Result<()> {
364    let status = header
365        .as_ref()
366        .and_then(|header| header.status.as_ref())
367        .context(IllegalDatabaseResponseSnafu {
368            err_msg: "either response header or status is missing",
369        })?;
370
371    if StatusCode::is_success(status.status_code) {
372        Ok(())
373    } else {
374        let code =
375            StatusCode::from_u32(status.status_code).context(IllegalDatabaseResponseSnafu {
376                err_msg: format!("unknown server status: {:?}", status),
377            })?;
378        ServerSnafu {
379            code,
380            msg: status.err_msg.clone(),
381        }
382        .fail()
383    }
384}
385
386#[cfg(test)]
387#[allow(deprecated)]
388mod test {
389    use api::v1::Status as PbStatus;
390    use api::v1::region::region_server::{Region, RegionServer};
391    use api::v1::region::{
392        BulkInsertRequest, RegionResponse as PbRegionResponse, RemoteDynFilterUnregister,
393        RemoteDynFilterUpdate, region_request, remote_dyn_filter_request,
394    };
395    use common_recordbatch::adapter::RecordBatchMetrics;
396    use datatypes::arrow::array::Int32Array;
397    use datatypes::prelude::{ConcreteDataType, VectorRef};
398    use datatypes::schema::{ColumnSchema, Schema};
399    use datatypes::vectors::Int32Vector;
400    use futures_util::stream;
401    use tokio::net::TcpListener;
402    use tokio_stream::wrappers::TcpListenerStream;
403    use tonic::codec::CompressionEncoding;
404    use tonic::{Request, Response, Status};
405
406    use super::*;
407    use crate::Error::{self, IllegalDatabaseResponse, Server};
408
409    #[derive(Clone)]
410    struct CompressionRecordingRegionService {
411        zstd_headers: Arc<std::sync::Mutex<Vec<(bool, bool)>>>,
412    }
413
414    #[tonic::async_trait]
415    impl Region for CompressionRecordingRegionService {
416        async fn handle(
417            &self,
418            request: Request<RegionRequest>,
419        ) -> std::result::Result<Response<PbRegionResponse>, Status> {
420            let metadata = request.metadata();
421            let sends_zstd = metadata
422                .get("grpc-encoding")
423                .and_then(|v| v.to_str().ok())
424                .is_some_and(|value| value == "zstd");
425            let accepts_zstd = metadata
426                .get("grpc-accept-encoding")
427                .and_then(|value| value.to_str().ok())
428                .is_some_and(|value| value.split(',').any(|encoding| encoding == "zstd"));
429            self.zstd_headers
430                .lock()
431                .unwrap()
432                .push((sends_zstd, accepts_zstd));
433
434            Ok(Response::new(PbRegionResponse {
435                header: Some(ResponseHeader {
436                    status: Some(PbStatus {
437                        status_code: StatusCode::Success as u32,
438                        ..Default::default()
439                    }),
440                }),
441                ..Default::default()
442            }))
443        }
444    }
445
446    #[tokio::test]
447    async fn test_inserts_and_bulk_insert_use_transport_compression() {
448        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
449        let addr = listener.local_addr().unwrap();
450        let zstd_headers = Arc::new(std::sync::Mutex::new(Vec::new()));
451        let service = CompressionRecordingRegionService {
452            zstd_headers: zstd_headers.clone(),
453        };
454        let server = tokio::spawn(async move {
455            tonic::transport::Server::builder()
456                .add_service(
457                    RegionServer::new(service)
458                        .accept_compressed(CompressionEncoding::Zstd)
459                        .send_compressed(CompressionEncoding::Zstd),
460                )
461                .serve_with_incoming(TcpListenerStream::new(listener))
462                .await
463                .unwrap();
464        });
465        let client = Client::with_urls([addr.to_string()]);
466        let requester = RegionRequester::new(client.clone(), true, true);
467        let send_only_requester = RegionRequester::new(client.clone(), true, false);
468        let accept_only_requester = RegionRequester::new(client.clone(), false, true);
469        let disabled_requester = RegionRequester::new(client, false, false);
470        let inserts = || RegionRequest {
471            body: Some(region_request::Body::Inserts(Default::default())),
472            ..Default::default()
473        };
474        let bulk_insert = || RegionRequest {
475            body: Some(region_request::Body::BulkInsert(
476                BulkInsertRequest::default(),
477            )),
478            ..Default::default()
479        };
480
481        requester.handle(inserts()).await.unwrap();
482        send_only_requester.handle(inserts()).await.unwrap();
483        accept_only_requester.handle(inserts()).await.unwrap();
484        requester.handle(bulk_insert()).await.unwrap();
485        disabled_requester.handle(bulk_insert()).await.unwrap();
486        requester
487            .handle(build_remote_dyn_filter_unregister_request(
488                "query-1",
489                RemoteDynFilterUnregister {
490                    filter_id: "filter-1".to_string(),
491                },
492            ))
493            .await
494            .unwrap();
495
496        assert_eq!(
497            vec![
498                (true, true),
499                (true, false),
500                (false, true),
501                (true, true),
502                (false, false),
503                (false, false)
504            ],
505            *zstd_headers.lock().unwrap()
506        );
507        server.abort();
508    }
509
510    #[test]
511    fn test_flight_stream_error_preserves_peer_address() {
512        let error = flight_stream_error(
513            "127.0.0.1:4001",
514            tonic::Status::unavailable("datanode unavailable").into(),
515        );
516
517        assert!(matches!(
518            error,
519            error::Error::FlightGet {
520                addr,
521                tonic_code: tonic::Code::Unavailable,
522                ..
523            } if addr == "127.0.0.1:4001"
524        ));
525    }
526
527    #[tokio::test]
528    async fn test_empty_flight_stream_preserves_peer_address() {
529        let Err(error) = recordbatches_from_flight_message_stream(
530            "127.0.0.1:4001".to_string(),
531            stream::empty::<Result<FlightMessage>>(),
532        )
533        .await
534        else {
535            panic!("expected empty Flight stream to fail");
536        };
537
538        assert!(matches!(
539            error,
540            error::Error::FlightGet {
541                addr,
542                tonic_code: tonic::Code::Unknown,
543                ..
544            } if addr == "127.0.0.1:4001"
545        ));
546    }
547
548    fn test_schema() -> Arc<Schema> {
549        Arc::new(Schema::new(vec![ColumnSchema::new(
550            "v",
551            ConcreteDataType::int32_datatype(),
552            false,
553        )]))
554    }
555
556    fn test_metrics_json() -> String {
557        serde_json::to_string(&RecordBatchMetrics {
558            elapsed_compute: 7,
559            ..Default::default()
560        })
561        .unwrap()
562    }
563
564    #[test]
565    fn test_check_response_header() {
566        let result = check_response_header(&None);
567        assert!(matches!(
568            result.unwrap_err(),
569            IllegalDatabaseResponse { .. }
570        ));
571
572        let result = check_response_header(&Some(ResponseHeader { status: None }));
573        assert!(matches!(
574            result.unwrap_err(),
575            IllegalDatabaseResponse { .. }
576        ));
577
578        let result = check_response_header(&Some(ResponseHeader {
579            status: Some(PbStatus {
580                status_code: StatusCode::Success as u32,
581                err_msg: String::default(),
582            }),
583        }));
584        assert!(result.is_ok());
585
586        let result = check_response_header(&Some(ResponseHeader {
587            status: Some(PbStatus {
588                status_code: u32::MAX,
589                err_msg: String::default(),
590            }),
591        }));
592        assert!(matches!(
593            result.unwrap_err(),
594            IllegalDatabaseResponse { .. }
595        ));
596
597        let result = check_response_header(&Some(ResponseHeader {
598            status: Some(PbStatus {
599                status_code: StatusCode::Internal as u32,
600                err_msg: "blabla".to_string(),
601            }),
602        }));
603        let Server { code, msg, .. } = result.unwrap_err() else {
604            unreachable!()
605        };
606        assert_eq!(code, StatusCode::Internal);
607        assert_eq!(msg, "blabla");
608    }
609
610    #[test]
611    fn test_build_remote_dyn_filter_request_sets_header_and_body() {
612        let request = build_remote_dyn_filter_update_request(
613            "query-1",
614            RemoteDynFilterUpdate {
615                filter_id: "filter-1".to_string(),
616                payload: vec![1, 2, 3],
617                generation: 7,
618                is_complete: false,
619            },
620        );
621
622        request.header.expect("remote dyn filter header must exist");
623
624        let body = request.body.expect("remote dyn filter body must exist");
625        let region_request::Body::RemoteDynFilter(remote_request) = body else {
626            panic!("expected remote dyn filter request body");
627        };
628
629        assert_eq!(remote_request.query_id, "query-1");
630        assert!(matches!(
631            remote_request.action,
632            Some(remote_dyn_filter_request::Action::Update(_))
633        ));
634    }
635
636    #[test]
637    fn test_build_remote_dyn_filter_unregister_request_sets_header_and_body() {
638        let request = build_remote_dyn_filter_unregister_request(
639            "query-1",
640            RemoteDynFilterUnregister {
641                filter_id: "filter-9".to_string(),
642            },
643        );
644
645        request.header.expect("remote dyn filter header must exist");
646
647        let body = request.body.expect("remote dyn filter body must exist");
648        let region_request::Body::RemoteDynFilter(remote_request) = body else {
649            panic!("expected remote dyn filter request body");
650        };
651
652        assert_eq!(remote_request.query_id, "query-1");
653        assert!(matches!(
654            remote_request.action,
655            Some(remote_dyn_filter_request::Action::Unregister(_))
656        ));
657    }
658
659    #[tokio::test]
660    async fn test_record_batch_stream_continues_after_pre_batch_metrics() {
661        let schema = test_schema();
662        let batch = RecordBatch::new(
663            schema.clone(),
664            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
665        )
666        .unwrap();
667
668        let mut recordbatches = recordbatches_from_flight_message_stream(
669            "test-peer".to_string(),
670            stream::iter(vec![
671                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
672                Ok(FlightMessage::Metrics(test_metrics_json())),
673                Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
674            ]),
675        )
676        .await
677        .unwrap();
678
679        let batch = recordbatches.next().await.unwrap().unwrap();
680        assert_eq!(batch.num_rows(), 1);
681        assert!(recordbatches.next().await.is_none());
682
683        let metrics = recordbatches.metrics().unwrap();
684        assert_eq!(metrics.elapsed_compute, 7);
685    }
686
687    #[tokio::test]
688    async fn test_record_batch_is_yielded_without_waiting_for_next_message() {
689        let schema = test_schema();
690        let batch = RecordBatch::new(
691            schema.clone(),
692            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
693        )
694        .unwrap();
695
696        let messages = stream::iter(vec![
697            Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
698            Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
699        ])
700        .chain(stream::pending::<Result<FlightMessage>>());
701        let mut recordbatches =
702            recordbatches_from_flight_message_stream("test-peer".to_string(), messages)
703                .await
704                .unwrap();
705
706        let batch = tokio::time::timeout(Duration::from_secs(1), recordbatches.next())
707            .await
708            .expect("the first batch must not wait for lookahead")
709            .unwrap()
710            .unwrap();
711        assert_eq!(batch.num_rows(), 1);
712    }
713
714    #[tokio::test]
715    async fn test_malformed_region_metrics_are_non_fatal() {
716        let schema = test_schema();
717        let batch = RecordBatch::new(
718            schema.clone(),
719            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
720        )
721        .unwrap();
722        let mut recordbatches = recordbatches_from_flight_message_stream(
723            "test-peer".to_string(),
724            stream::iter(vec![
725                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
726                Ok(FlightMessage::Metrics("{not-json}".to_string())),
727                Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
728            ]),
729        )
730        .await
731        .unwrap();
732
733        assert_eq!(recordbatches.next().await.unwrap().unwrap().num_rows(), 1);
734        assert!(recordbatches.next().await.is_none());
735    }
736
737    #[tokio::test]
738    async fn test_record_batch_stream_preserves_following_record_batch() {
739        let schema = test_schema();
740        let first_batch = RecordBatch::new(
741            schema.clone(),
742            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
743        )
744        .unwrap();
745        let second_batch = RecordBatch::new(
746            schema.clone(),
747            vec![Arc::new(Int32Vector::from_slice([2])) as VectorRef],
748        )
749        .unwrap();
750
751        let mut recordbatches = recordbatches_from_flight_message_stream(
752            "test-peer".to_string(),
753            stream::iter(vec![
754                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
755                Ok(FlightMessage::RecordBatch(
756                    first_batch.into_df_record_batch(),
757                )),
758                Ok(FlightMessage::RecordBatch(
759                    second_batch.into_df_record_batch(),
760                )),
761            ]),
762        )
763        .await
764        .unwrap();
765
766        let first_batch = recordbatches.next().await.unwrap().unwrap();
767        let second_batch = recordbatches.next().await.unwrap().unwrap();
768        assert_eq!(
769            first_batch
770                .column(0)
771                .as_any()
772                .downcast_ref::<Int32Array>()
773                .unwrap()
774                .value(0),
775            1
776        );
777        assert_eq!(
778            second_batch
779                .column(0)
780                .as_any()
781                .downcast_ref::<Int32Array>()
782                .unwrap()
783                .value(0),
784            2
785        );
786        assert!(recordbatches.next().await.is_none());
787    }
788
789    #[tokio::test]
790    async fn test_record_batch_stream_captures_final_metrics_after_record_batch() {
791        let schema = test_schema();
792        let batch = RecordBatch::new(
793            schema.clone(),
794            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
795        )
796        .unwrap();
797        let final_metrics = serde_json::to_string(&RecordBatchMetrics {
798            elapsed_compute: 99,
799            ..Default::default()
800        })
801        .unwrap();
802        let mut recordbatches = recordbatches_from_flight_message_stream(
803            "test-peer".to_string(),
804            stream::iter(vec![
805                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
806                Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
807                Ok(FlightMessage::Metrics(final_metrics)),
808            ]),
809        )
810        .await
811        .unwrap();
812
813        assert_eq!(recordbatches.next().await.unwrap().unwrap().num_rows(), 1);
814        assert!(recordbatches.next().await.is_none());
815        assert_eq!(recordbatches.metrics().unwrap().elapsed_compute, 99);
816    }
817
818    #[tokio::test]
819    async fn test_record_batch_stream_exposes_error_after_pre_batch_metrics() {
820        let schema = test_schema();
821        let mut recordbatches = recordbatches_from_flight_message_stream(
822            "test-peer".to_string(),
823            stream::iter(vec![
824                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
825                Ok(FlightMessage::Metrics(test_metrics_json())),
826                Err(Error::from(Status::internal("boom after metrics"))),
827            ]),
828        )
829        .await
830        .unwrap();
831
832        let err = recordbatches.next().await.unwrap().unwrap_err();
833        assert_eq!("External error", err.to_string());
834        assert!(
835            format!("{err:?}").contains("boom after metrics"),
836            "unexpected error: {err:?}"
837        );
838        assert!(recordbatches.next().await.is_none());
839    }
840}