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::{FlightMessageKind, 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        let mut stream_ended = false;
251
252        while !stream_ended {
253            let flight_message = match reader.read_next().await {
254                Ok(Some(message)) => message,
255                Ok(None) => break,
256                Err(error) => {
257                    yield Err(BoxedError::new(flight_stream_error(&stream_addr, error)))
258                        .context(ExternalSnafu);
259                    break;
260                }
261            };
262
263            match flight_message {
264                FlightMessage::RecordBatch(record_batch) => {
265                    let result_to_yield =
266                        RecordBatch::from_df_record_batch(schema_cloned.clone(), record_batch);
267
268                    // Metrics follow a batch so MergeScan can observe them before yielding it.
269                    match reader.peek_next_message_kind().await {
270                        Ok(Some(FlightMessageKind::Metrics)) => {
271                            let metrics_message = match reader.read_next().await {
272                                Ok(Some(FlightMessage::Metrics(metrics))) => metrics,
273                                Ok(Some(_) | None) => {
274                                    yield IllegalFlightMessagesSnafu {
275                                        reason: "Flight stream changed after peek",
276                                    }
277                                    .fail()
278                                    .map_err(BoxedError::new)
279                                    .context(ExternalSnafu);
280                                    break;
281                                }
282                                Err(error) => {
283                                    yield Err(BoxedError::new(flight_stream_error(
284                                        &stream_addr,
285                                        error,
286                                    )))
287                                    .context(ExternalSnafu);
288                                    break;
289                                }
290                            };
291                            let metrics = serde_json::from_str(&metrics_message).ok().map(Arc::new);
292                            metrics_ref.swap(metrics);
293                        }
294                        Ok(Some(FlightMessageKind::RecordBatch)) => {}
295                        Ok(Some(FlightMessageKind::Schema | FlightMessageKind::AffectedRows)) => {
296                            yield IllegalFlightMessagesSnafu {
297                                reason: "A RecordBatch message can only be succeeded by a Metrics message or another RecordBatch message"
298                            }
299                            .fail()
300                            .map_err(BoxedError::new)
301                            .context(ExternalSnafu);
302                            break;
303                        }
304                        Ok(None) => stream_ended = true,
305                        Err(error) => {
306                            yield Err(BoxedError::new(flight_stream_error(&stream_addr, error)))
307                                .context(ExternalSnafu);
308                            break;
309                        }
310                    }
311
312                    yield Ok(result_to_yield);
313                }
314                FlightMessage::Metrics(s) => {
315                    // Metrics may arrive before the next RecordBatch.
316                    let m = serde_json::from_str(&s).ok().map(Arc::new);
317                    metrics_ref.swap(m);
318                    continue;
319                }
320                _ => {
321                    yield IllegalFlightMessagesSnafu {
322                        reason: "A Schema message must be succeeded exclusively by a set of RecordBatch messages"
323                    }
324                    .fail()
325                    .map_err(BoxedError::new)
326                    .context(ExternalSnafu);
327                    break;
328                }
329            }
330        }
331    }));
332    let record_batch_stream = RecordBatchStreamWrapper {
333        schema,
334        stream,
335        output_ordering: None,
336        metrics,
337        span: Span::current(),
338    };
339    Ok(Box::pin(record_batch_stream))
340}
341
342fn flight_stream_error(addr: &str, error: error::Error) -> error::Error {
343    let tonic_code = error.tonic_code().unwrap_or(tonic::Code::Unknown);
344    if error.status_code().should_log_error() {
345        error!(
346            error; "Failed to receive Flight data, addr: {}, code: {}",
347            addr,
348            tonic_code
349        );
350    }
351
352    error::Error::FlightGet {
353        addr: addr.to_string(),
354        tonic_code,
355        source: BoxedError::new(error),
356    }
357}
358
359pub fn build_remote_dyn_filter_update_request(
360    query_id: impl Into<String>,
361    update: RemoteDynFilterUpdate,
362) -> RegionRequest {
363    build_remote_dyn_filter_request(
364        query_id.into(),
365        remote_dyn_filter_request::Action::Update(update),
366    )
367}
368
369pub fn build_remote_dyn_filter_unregister_request(
370    query_id: impl Into<String>,
371    unregister: RemoteDynFilterUnregister,
372) -> RegionRequest {
373    build_remote_dyn_filter_request(
374        query_id.into(),
375        remote_dyn_filter_request::Action::Unregister(unregister),
376    )
377}
378
379fn build_remote_dyn_filter_request(
380    query_id: String,
381    action: remote_dyn_filter_request::Action,
382) -> RegionRequest {
383    RegionRequest {
384        header: Some(RegionRequestHeader {
385            tracing_context: TracingContext::from_current_span().to_w3c(),
386            ..Default::default()
387        }),
388        body: Some(region_request::Body::RemoteDynFilter(
389            RemoteDynFilterRequest {
390                query_id,
391                action: Some(action),
392            },
393        )),
394    }
395}
396
397pub fn check_response_header(header: &Option<ResponseHeader>) -> Result<()> {
398    let status = header
399        .as_ref()
400        .and_then(|header| header.status.as_ref())
401        .context(IllegalDatabaseResponseSnafu {
402            err_msg: "either response header or status is missing",
403        })?;
404
405    if StatusCode::is_success(status.status_code) {
406        Ok(())
407    } else {
408        let code =
409            StatusCode::from_u32(status.status_code).context(IllegalDatabaseResponseSnafu {
410                err_msg: format!("unknown server status: {:?}", status),
411            })?;
412        ServerSnafu {
413            code,
414            msg: status.err_msg.clone(),
415        }
416        .fail()
417    }
418}
419
420#[cfg(test)]
421mod test {
422    use api::v1::Status as PbStatus;
423    use api::v1::region::region_server::{Region, RegionServer};
424    use api::v1::region::{
425        BulkInsertRequest, RegionResponse as PbRegionResponse, RemoteDynFilterUnregister,
426        RemoteDynFilterUpdate, region_request, remote_dyn_filter_request,
427    };
428    use common_recordbatch::adapter::RecordBatchMetrics;
429    use datatypes::arrow::array::Int32Array;
430    use datatypes::prelude::{ConcreteDataType, VectorRef};
431    use datatypes::schema::{ColumnSchema, Schema};
432    use datatypes::vectors::Int32Vector;
433    use futures_util::stream;
434    use tokio::net::TcpListener;
435    use tokio_stream::wrappers::TcpListenerStream;
436    use tonic::codec::CompressionEncoding;
437    use tonic::{Request, Response, Status};
438
439    use super::*;
440    use crate::Error::{self, IllegalDatabaseResponse, Server};
441
442    #[derive(Clone)]
443    struct CompressionRecordingRegionService {
444        zstd_headers: Arc<std::sync::Mutex<Vec<(bool, bool)>>>,
445    }
446
447    #[tonic::async_trait]
448    impl Region for CompressionRecordingRegionService {
449        async fn handle(
450            &self,
451            request: Request<RegionRequest>,
452        ) -> std::result::Result<Response<PbRegionResponse>, Status> {
453            let metadata = request.metadata();
454            let sends_zstd = metadata
455                .get("grpc-encoding")
456                .and_then(|v| v.to_str().ok())
457                .is_some_and(|value| value == "zstd");
458            let accepts_zstd = metadata
459                .get("grpc-accept-encoding")
460                .and_then(|value| value.to_str().ok())
461                .is_some_and(|value| value.split(',').any(|encoding| encoding == "zstd"));
462            self.zstd_headers
463                .lock()
464                .unwrap()
465                .push((sends_zstd, accepts_zstd));
466
467            Ok(Response::new(PbRegionResponse {
468                header: Some(ResponseHeader {
469                    status: Some(PbStatus {
470                        status_code: StatusCode::Success as u32,
471                        ..Default::default()
472                    }),
473                }),
474                ..Default::default()
475            }))
476        }
477    }
478
479    #[tokio::test]
480    async fn test_inserts_and_bulk_insert_use_transport_compression() {
481        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
482        let addr = listener.local_addr().unwrap();
483        let zstd_headers = Arc::new(std::sync::Mutex::new(Vec::new()));
484        let service = CompressionRecordingRegionService {
485            zstd_headers: zstd_headers.clone(),
486        };
487        let server = tokio::spawn(async move {
488            tonic::transport::Server::builder()
489                .add_service(
490                    RegionServer::new(service)
491                        .accept_compressed(CompressionEncoding::Zstd)
492                        .send_compressed(CompressionEncoding::Zstd),
493                )
494                .serve_with_incoming(TcpListenerStream::new(listener))
495                .await
496                .unwrap();
497        });
498        let client = Client::with_urls([addr.to_string()]);
499        let requester = RegionRequester::new(client.clone(), true, true);
500        let send_only_requester = RegionRequester::new(client.clone(), true, false);
501        let accept_only_requester = RegionRequester::new(client.clone(), false, true);
502        let disabled_requester = RegionRequester::new(client, false, false);
503        let inserts = || RegionRequest {
504            body: Some(region_request::Body::Inserts(Default::default())),
505            ..Default::default()
506        };
507        let bulk_insert = || RegionRequest {
508            body: Some(region_request::Body::BulkInsert(
509                BulkInsertRequest::default(),
510            )),
511            ..Default::default()
512        };
513
514        requester.handle(inserts()).await.unwrap();
515        send_only_requester.handle(inserts()).await.unwrap();
516        accept_only_requester.handle(inserts()).await.unwrap();
517        requester.handle(bulk_insert()).await.unwrap();
518        disabled_requester.handle(bulk_insert()).await.unwrap();
519        requester
520            .handle(build_remote_dyn_filter_unregister_request(
521                "query-1",
522                RemoteDynFilterUnregister {
523                    filter_id: "filter-1".to_string(),
524                },
525            ))
526            .await
527            .unwrap();
528
529        assert_eq!(
530            vec![
531                (true, true),
532                (true, false),
533                (false, true),
534                (true, true),
535                (false, false),
536                (false, false)
537            ],
538            *zstd_headers.lock().unwrap()
539        );
540        server.abort();
541    }
542
543    #[test]
544    fn test_flight_stream_error_preserves_peer_address() {
545        let error = flight_stream_error(
546            "127.0.0.1:4001",
547            tonic::Status::unavailable("datanode unavailable").into(),
548        );
549
550        assert!(matches!(
551            error,
552            error::Error::FlightGet {
553                addr,
554                tonic_code: tonic::Code::Unavailable,
555                ..
556            } if addr == "127.0.0.1:4001"
557        ));
558    }
559
560    #[tokio::test]
561    async fn test_empty_flight_stream_preserves_peer_address() {
562        let Err(error) = recordbatches_from_flight_message_stream(
563            "127.0.0.1:4001".to_string(),
564            stream::empty::<Result<FlightMessage>>(),
565        )
566        .await
567        else {
568            panic!("expected empty Flight stream to fail");
569        };
570
571        assert!(matches!(
572            error,
573            error::Error::FlightGet {
574                addr,
575                tonic_code: tonic::Code::Unknown,
576                ..
577            } if addr == "127.0.0.1:4001"
578        ));
579    }
580
581    fn test_schema() -> Arc<Schema> {
582        Arc::new(Schema::new(vec![ColumnSchema::new(
583            "v",
584            ConcreteDataType::int32_datatype(),
585            false,
586        )]))
587    }
588
589    fn test_metrics_json() -> String {
590        serde_json::to_string(&RecordBatchMetrics {
591            elapsed_compute: 7,
592            ..Default::default()
593        })
594        .unwrap()
595    }
596
597    #[test]
598    fn test_check_response_header() {
599        let result = check_response_header(&None);
600        assert!(matches!(
601            result.unwrap_err(),
602            IllegalDatabaseResponse { .. }
603        ));
604
605        let result = check_response_header(&Some(ResponseHeader { status: None }));
606        assert!(matches!(
607            result.unwrap_err(),
608            IllegalDatabaseResponse { .. }
609        ));
610
611        let result = check_response_header(&Some(ResponseHeader {
612            status: Some(PbStatus {
613                status_code: StatusCode::Success as u32,
614                err_msg: String::default(),
615            }),
616        }));
617        assert!(result.is_ok());
618
619        let result = check_response_header(&Some(ResponseHeader {
620            status: Some(PbStatus {
621                status_code: u32::MAX,
622                err_msg: String::default(),
623            }),
624        }));
625        assert!(matches!(
626            result.unwrap_err(),
627            IllegalDatabaseResponse { .. }
628        ));
629
630        let result = check_response_header(&Some(ResponseHeader {
631            status: Some(PbStatus {
632                status_code: StatusCode::Internal as u32,
633                err_msg: "blabla".to_string(),
634            }),
635        }));
636        let Server { code, msg, .. } = result.unwrap_err() else {
637            unreachable!()
638        };
639        assert_eq!(code, StatusCode::Internal);
640        assert_eq!(msg, "blabla");
641    }
642
643    #[test]
644    fn test_build_remote_dyn_filter_request_sets_header_and_body() {
645        let request = build_remote_dyn_filter_update_request(
646            "query-1",
647            RemoteDynFilterUpdate {
648                filter_id: "filter-1".to_string(),
649                payload: vec![1, 2, 3],
650                generation: 7,
651                is_complete: false,
652            },
653        );
654
655        request.header.expect("remote dyn filter header must exist");
656
657        let body = request.body.expect("remote dyn filter body must exist");
658        let region_request::Body::RemoteDynFilter(remote_request) = body else {
659            panic!("expected remote dyn filter request body");
660        };
661
662        assert_eq!(remote_request.query_id, "query-1");
663        assert!(matches!(
664            remote_request.action,
665            Some(remote_dyn_filter_request::Action::Update(_))
666        ));
667    }
668
669    #[test]
670    fn test_build_remote_dyn_filter_unregister_request_sets_header_and_body() {
671        let request = build_remote_dyn_filter_unregister_request(
672            "query-1",
673            RemoteDynFilterUnregister {
674                filter_id: "filter-9".to_string(),
675            },
676        );
677
678        request.header.expect("remote dyn filter header must exist");
679
680        let body = request.body.expect("remote dyn filter body must exist");
681        let region_request::Body::RemoteDynFilter(remote_request) = body else {
682            panic!("expected remote dyn filter request body");
683        };
684
685        assert_eq!(remote_request.query_id, "query-1");
686        assert!(matches!(
687            remote_request.action,
688            Some(remote_dyn_filter_request::Action::Unregister(_))
689        ));
690    }
691
692    #[tokio::test]
693    async fn test_record_batch_stream_continues_after_pre_batch_metrics() {
694        let schema = test_schema();
695        let batch = RecordBatch::new(
696            schema.clone(),
697            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
698        )
699        .unwrap();
700
701        let mut recordbatches = recordbatches_from_flight_message_stream(
702            "test-peer".to_string(),
703            stream::iter(vec![
704                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
705                Ok(FlightMessage::Metrics(test_metrics_json())),
706                Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
707            ]),
708        )
709        .await
710        .unwrap();
711
712        let batch = recordbatches.next().await.unwrap().unwrap();
713        assert_eq!(batch.num_rows(), 1);
714        assert!(recordbatches.next().await.is_none());
715
716        let metrics = recordbatches.metrics().unwrap();
717        assert_eq!(metrics.elapsed_compute, 7);
718    }
719
720    #[tokio::test]
721    async fn test_record_batch_stream_updates_following_metrics_before_yielding_batch() {
722        let schema = test_schema();
723        let batch = RecordBatch::new(
724            schema.clone(),
725            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
726        )
727        .unwrap();
728        let mut recordbatches = recordbatches_from_flight_message_stream(
729            "test-peer".to_string(),
730            stream::iter(vec![
731                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
732                Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
733                Ok(FlightMessage::Metrics(test_metrics_json())),
734            ]),
735        )
736        .await
737        .unwrap();
738
739        let batch = recordbatches.next().await.unwrap().unwrap();
740        assert_eq!(batch.num_rows(), 1);
741
742        let metrics = recordbatches.metrics().unwrap();
743        assert_eq!(metrics.elapsed_compute, 7);
744        assert!(recordbatches.next().await.is_none());
745    }
746
747    #[tokio::test]
748    async fn test_record_batch_stream_preserves_peeked_record_batch() {
749        let schema = test_schema();
750        let first_batch = RecordBatch::new(
751            schema.clone(),
752            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
753        )
754        .unwrap();
755        let second_batch = RecordBatch::new(
756            schema.clone(),
757            vec![Arc::new(Int32Vector::from_slice([2])) as VectorRef],
758        )
759        .unwrap();
760
761        let mut recordbatches = recordbatches_from_flight_message_stream(
762            "test-peer".to_string(),
763            stream::iter(vec![
764                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
765                Ok(FlightMessage::RecordBatch(
766                    first_batch.into_df_record_batch(),
767                )),
768                Ok(FlightMessage::RecordBatch(
769                    second_batch.into_df_record_batch(),
770                )),
771            ]),
772        )
773        .await
774        .unwrap();
775
776        let first_batch = recordbatches.next().await.unwrap().unwrap();
777        let second_batch = recordbatches.next().await.unwrap().unwrap();
778        assert_eq!(
779            first_batch
780                .column(0)
781                .as_any()
782                .downcast_ref::<Int32Array>()
783                .unwrap()
784                .value(0),
785            1
786        );
787        assert_eq!(
788            second_batch
789                .column(0)
790                .as_any()
791                .downcast_ref::<Int32Array>()
792                .unwrap()
793                .value(0),
794            2
795        );
796        assert!(recordbatches.next().await.is_none());
797    }
798
799    #[tokio::test]
800    async fn test_record_batch_stream_captures_final_metrics_after_record_batch() {
801        let schema = test_schema();
802        let batch = RecordBatch::new(
803            schema.clone(),
804            vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
805        )
806        .unwrap();
807        let final_metrics = serde_json::to_string(&RecordBatchMetrics {
808            elapsed_compute: 99,
809            ..Default::default()
810        })
811        .unwrap();
812        let mut recordbatches = recordbatches_from_flight_message_stream(
813            "test-peer".to_string(),
814            stream::iter(vec![
815                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
816                Ok(FlightMessage::RecordBatch(batch.into_df_record_batch())),
817                Ok(FlightMessage::Metrics(final_metrics)),
818            ]),
819        )
820        .await
821        .unwrap();
822
823        assert_eq!(recordbatches.next().await.unwrap().unwrap().num_rows(), 1);
824        assert!(recordbatches.next().await.is_none());
825        assert_eq!(recordbatches.metrics().unwrap().elapsed_compute, 99);
826    }
827
828    #[tokio::test]
829    async fn test_record_batch_stream_exposes_error_after_pre_batch_metrics() {
830        let schema = test_schema();
831        let mut recordbatches = recordbatches_from_flight_message_stream(
832            "test-peer".to_string(),
833            stream::iter(vec![
834                Ok(FlightMessage::Schema(schema.arrow_schema().clone())),
835                Ok(FlightMessage::Metrics(test_metrics_json())),
836                Err(Error::from(Status::internal("boom after metrics"))),
837            ]),
838        )
839        .await
840        .unwrap();
841
842        let err = recordbatches.next().await.unwrap().unwrap_err();
843        assert_eq!("External error", err.to_string());
844        assert!(
845            format!("{err:?}").contains("boom after metrics"),
846            "unexpected error: {err:?}"
847        );
848        assert!(recordbatches.next().await.is_none());
849    }
850}