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::{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 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 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 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 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}