1use 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 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 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 yield Ok(RecordBatch::from_df_record_batch(
267 schema_cloned.clone(),
268 record_batch,
269 ));
270 }
271 FlightMessage::Metrics(s) => {
272 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}