1use std::collections::VecDeque;
16use std::future::Future;
17use std::pin::Pin;
18use std::task::{Context, Poll};
19use std::time::{Duration, Instant};
20
21use arrow_flight::FlightData;
22use common_error::ext::ErrorExt;
23use common_grpc::flight::{FlightEncoder, FlightMessage};
24use common_recordbatch::recordbatch::merge_record_batches;
25use common_recordbatch::{RecordBatch, SendableRecordBatchStream};
26use common_telemetry::tracing::{Instrument, info_span};
27use common_telemetry::tracing_context::{FutureExt, TracingContext};
28use common_telemetry::{error, info, warn};
29use datatypes::schema::SchemaRef;
30use futures::channel::mpsc;
31use futures::channel::mpsc::Sender;
32use futures::future::poll_fn;
33use futures::{SinkExt, Stream, StreamExt};
34use pin_project::{pin_project, pinned_drop};
35use session::context::{
36 FLIGHT_METRICS_HEARTBEAT_INTERVAL, QueryContextRef,
37 SUPPORT_FLIGHT_METRICS_BEFORE_BATCH_EXTENSION_KEY,
38};
39use snafu::ResultExt;
40use tokio::task::JoinHandle;
41use tokio::time;
42
43use crate::error;
44use crate::grpc::FlightCompression;
45use crate::grpc::flight::TonicResult;
46
47pub enum FlightRecordBatchSource {
48 RecordBatches(SendableRecordBatchStream),
49 AffectedRows {
50 rows: usize,
51 metrics: Option<String>,
52 },
53}
54
55pub enum FlightRecordBatchStreamInput<F = std::future::Ready<TonicResult<FlightRecordBatchSource>>>
57{
58 Ready(FlightRecordBatchSource),
59 Initializer(F),
60}
61
62impl FlightRecordBatchStreamInput {
63 pub fn ready(source: FlightRecordBatchSource) -> Self {
65 Self::Ready(source)
66 }
67}
68
69impl<F> FlightRecordBatchStreamInput<F> {
70 pub fn initializer(initializer: F) -> Self {
74 Self::Initializer(initializer)
75 }
76}
77
78struct StreamMetrics {
80 send_schema_duration: Duration,
81 send_record_batch_duration: Duration,
82 send_metrics_duration: Duration,
83 fetch_content_duration: Duration,
84 record_batch_count: usize,
85 metrics_count: usize,
86 total_rows: usize,
87 total_bytes: usize,
88 should_log: bool,
89}
90
91impl StreamMetrics {
92 fn new(should_log: bool) -> Self {
93 Self {
94 send_schema_duration: Duration::ZERO,
95 send_record_batch_duration: Duration::ZERO,
96 send_metrics_duration: Duration::ZERO,
97 fetch_content_duration: Duration::ZERO,
98 record_batch_count: 0,
99 metrics_count: 0,
100 total_rows: 0,
101 total_bytes: 0,
102 should_log,
103 }
104 }
105}
106
107impl Drop for StreamMetrics {
108 fn drop(&mut self) {
109 if self.should_log {
110 info!(
111 "flight_data_stream finished: \
112 send_schema_duration={:?}, \
113 send_record_batch_duration={:?}, \
114 send_metrics_duration={:?}, \
115 fetch_content_duration={:?}, \
116 record_batch_count={}, \
117 metrics_count={}, \
118 total_rows={}, \
119 total_bytes={}",
120 self.send_schema_duration,
121 self.send_record_batch_duration,
122 self.send_metrics_duration,
123 self.fetch_content_duration,
124 self.record_batch_count,
125 self.metrics_count,
126 self.total_rows,
127 self.total_bytes
128 );
129 }
130 }
131}
132
133struct BatchAccumulator {
141 batches: Vec<RecordBatch>,
142 rows: usize,
143 bytes: usize,
144}
145
146impl BatchAccumulator {
147 const MAX_ROWS: usize = 4096;
152 const MAX_BYTES: usize = 256 * 1024;
153 const MAX_BATCHES: usize = 16;
154
155 fn new() -> Self {
156 Self {
157 batches: Vec::new(),
158 rows: 0,
159 bytes: 0,
160 }
161 }
162
163 fn reaches_budget(rows: usize, bytes: usize) -> bool {
168 rows >= Self::MAX_ROWS || bytes >= Self::MAX_BYTES
169 }
170
171 fn is_empty(&self) -> bool {
172 self.batches.is_empty()
173 }
174
175 fn push(&mut self, batch: RecordBatch) -> bool {
177 self.rows += batch.num_rows();
178 self.bytes += batch.df_record_batch().get_array_memory_size();
179 self.batches.push(batch);
180 self.batches.len() >= Self::MAX_BATCHES || Self::reaches_budget(self.rows, self.bytes)
182 }
183
184 fn drain(&mut self) -> Vec<RecordBatch> {
188 self.rows = 0;
189 self.bytes = 0;
190 std::mem::take(&mut self.batches)
191 }
192}
193
194struct CoalescingBatcher {
197 acc: BatchAccumulator,
198 sent_first_batch: bool,
199 recordbatch_schema: SchemaRef,
200}
201
202impl CoalescingBatcher {
203 fn new(recordbatch_schema: SchemaRef) -> Self {
204 Self {
205 acc: BatchAccumulator::new(),
206 sent_first_batch: false,
207 recordbatch_schema,
208 }
209 }
210
211 async fn run(
215 &mut self,
216 recordbatches: &mut SendableRecordBatchStream,
217 tx: &mut Sender<TonicResult<FlightMessage>>,
218 metrics: &mut StreamMetrics,
219 ) -> bool {
220 loop {
221 let start = Instant::now();
222 let batch_or_err = recordbatches.next().in_current_span().await;
223 metrics.fetch_content_duration += start.elapsed();
224 let Some(batch_or_err) = batch_or_err else {
225 break;
226 };
227 let recordbatch = match batch_or_err {
228 Ok(recordbatch) => recordbatch,
229 Err(e) => {
230 if e.status_code().should_log_error() {
231 error!("{e:?}");
232 }
233 let e = Err(e).context(error::CollectRecordbatchSnafu);
234 if let Err(e) = tx.send(e.map_err(|x| x.into())).await {
235 warn!(e; "stop sending Flight data");
236 }
237 return false;
238 }
239 };
240 let batch_rows = recordbatch.num_rows();
241 let batch_bytes = recordbatch.df_record_batch().get_array_memory_size();
242 metrics.total_rows += batch_rows;
243 metrics.record_batch_count += 1;
244 metrics.total_bytes += batch_bytes;
245
246 if !self.sent_first_batch || BatchAccumulator::reaches_budget(batch_rows, batch_bytes) {
252 let start = Instant::now();
253 if let Err(e) = tx
254 .send(Ok(FlightMessage::RecordBatch(
255 recordbatch.into_df_record_batch(),
256 )))
257 .await
258 {
259 warn!(e; "stop sending Flight data");
260 return false;
261 }
262 metrics.send_record_batch_duration += start.elapsed();
263 self.sent_first_batch = true;
264 continue;
265 }
266
267 debug_assert!(self.acc.is_empty(), "the previous group must be flushed");
272 let mut should_flush = self.acc.push(recordbatch);
273 let mut eof = false;
274 let mut stream_error = None;
275 let mut pending_oversized = None;
276 while !should_flush {
277 let start = Instant::now();
278 let next = poll_fn(|cx| Poll::Ready(recordbatches.as_mut().poll_next(cx))).await;
279 metrics.fetch_content_duration += start.elapsed();
280 match next {
281 Poll::Ready(Some(Ok(recordbatch))) => {
282 let batch_rows = recordbatch.num_rows();
285 let batch_bytes = recordbatch.df_record_batch().get_array_memory_size();
286 metrics.total_rows += batch_rows;
287 metrics.record_batch_count += 1;
288 metrics.total_bytes += batch_bytes;
289 if BatchAccumulator::reaches_budget(batch_rows, batch_bytes) {
290 pending_oversized = Some(recordbatch);
295 break;
296 }
297 should_flush = self.acc.push(recordbatch);
298 }
299 Poll::Ready(Some(Err(e))) => {
300 stream_error = Some(e);
301 break;
302 }
303 Poll::Ready(None) => {
304 eof = true;
305 break;
306 }
307 Poll::Pending => break,
308 }
309 }
310
311 let mut batches = self.acc.drain();
312 if batches.len() >= 2
313 && let Ok(merged) = merge_record_batches(self.recordbatch_schema.clone(), &batches)
314 {
315 batches.clear();
319 batches.push(merged);
320 }
321 for recordbatch in batches {
322 let start = Instant::now();
323 if let Err(e) = tx
324 .send(Ok(FlightMessage::RecordBatch(
325 recordbatch.into_df_record_batch(),
326 )))
327 .await
328 {
329 warn!(e; "stop sending Flight data");
330 return false;
331 }
332 metrics.send_record_batch_duration += start.elapsed();
333 }
334 if let Some(recordbatch) = pending_oversized {
340 let start = Instant::now();
341 if let Err(e) = tx
342 .send(Ok(FlightMessage::RecordBatch(
343 recordbatch.into_df_record_batch(),
344 )))
345 .await
346 {
347 warn!(e; "stop sending Flight data");
348 return false;
349 }
350 metrics.send_record_batch_duration += start.elapsed();
351 }
352 if let Some(e) = stream_error {
353 if e.status_code().should_log_error() {
354 error!("{e:?}");
355 }
356 let e = Err(e).context(error::CollectRecordbatchSnafu);
357 if let Err(e) = tx.send(e.map_err(|x| x.into())).await {
358 warn!(e; "stop sending Flight data");
359 }
360 return false;
361 }
362 if eof {
363 break;
364 }
365 }
366 true
367 }
368}
369
370#[pin_project(PinnedDrop)]
371pub struct FlightRecordBatchStream {
372 #[pin]
373 rx: mpsc::Receiver<Result<FlightMessage, tonic::Status>>,
374 join_handle: JoinHandle<()>,
375 done: bool,
376 encoder: FlightEncoder,
377 buffer: VecDeque<FlightData>,
378}
379
380impl FlightRecordBatchStream {
381 async fn send_metrics(
382 tx: &mut Sender<TonicResult<FlightMessage>>,
383 metrics: &mut StreamMetrics,
384 metrics_str: String,
385 ) -> bool {
386 metrics.metrics_count += 1;
387 let start = Instant::now();
388 if let Err(e) = tx.send(Ok(FlightMessage::Metrics(metrics_str))).await {
389 warn!(e; "stop sending Flight data");
390 return false;
391 }
392 metrics.send_metrics_duration += start.elapsed();
393 true
394 }
395
396 async fn send_metrics_if_changed(
397 tx: &mut Sender<TonicResult<FlightMessage>>,
398 metrics: &mut StreamMetrics,
399 last_metrics_str: &mut Option<String>,
400 metrics_str: String,
401 ) -> bool {
402 if last_metrics_str.as_deref() == Some(metrics_str.as_str()) {
403 return true;
404 }
405
406 *last_metrics_str = Some(metrics_str.clone());
407 Self::send_metrics(tx, metrics, metrics_str).await
408 }
409
410 pub fn new<F>(
411 input: FlightRecordBatchStreamInput<F>,
412 tracing_context: TracingContext,
413 compression: FlightCompression,
414 query_ctx: QueryContextRef,
415 ) -> Self
416 where
417 F: Future<Output = TonicResult<FlightRecordBatchSource>> + Send + 'static,
418 {
419 let (mut tx, rx) = mpsc::channel::<TonicResult<FlightMessage>>(1);
420 let source_type = match &input {
421 FlightRecordBatchStreamInput::Ready(FlightRecordBatchSource::RecordBatches(_)) => {
422 "record_batches"
423 }
424 FlightRecordBatchStreamInput::Ready(FlightRecordBatchSource::AffectedRows {
425 ..
426 }) => "affected_rows",
427 FlightRecordBatchStreamInput::Initializer(_) => "initializer",
428 };
429 let initializer_tracing_context = tracing_context.clone();
430 let join_handle = common_runtime::spawn_global(
431 async move {
432 let source = async move {
433 match input {
434 FlightRecordBatchStreamInput::Ready(source) => Ok(source),
435 FlightRecordBatchStreamInput::Initializer(initializer) => initializer.await,
436 }
437 }
438 .trace(
439 initializer_tracing_context
440 .attach(info_span!("flight_data_stream_init", source_type)),
441 )
442 .await;
443
444 match source {
445 Ok(FlightRecordBatchSource::RecordBatches(recordbatches)) => {
446 let should_send_partial_metrics = query_ctx.explain_verbose();
448 let can_send_metrics_before_batch =
449 query_ctx.explain_verbose()
450 && query_ctx.live_analyze_metrics_enabled()
451 && query_ctx
452 .remote_query_id()
453 .zip(query_ctx.extension(
454 SUPPORT_FLIGHT_METRICS_BEFORE_BATCH_EXTENSION_KEY,
455 ))
456 .is_some_and(|(remote_query_id, capability)| {
457 capability == remote_query_id
458 });
459 Self::flight_data_stream(
460 recordbatches,
461 tx,
462 should_send_partial_metrics,
463 can_send_metrics_before_batch,
464 )
465 .await;
466 }
467 Ok(FlightRecordBatchSource::AffectedRows { rows, metrics }) => {
468 let _ = tx
469 .send(Ok(FlightMessage::AffectedRows { rows, metrics }))
470 .await;
471 }
472 Err(status) => {
473 let _ = tx.send(Err(status)).await;
474 }
475 }
476 }
477 .trace(tracing_context.attach(info_span!("flight_data_stream"))),
478 );
479 let encoder = if compression.arrow_compression() {
480 FlightEncoder::default()
481 } else {
482 FlightEncoder::with_compression_disabled()
483 };
484 Self {
485 rx,
486 join_handle,
487 done: false,
488 encoder,
489 buffer: VecDeque::new(),
490 }
491 }
492
493 async fn flight_data_stream(
494 mut recordbatches: SendableRecordBatchStream,
495 mut tx: Sender<TonicResult<FlightMessage>>,
496 should_send_partial_metrics: bool,
497 can_send_metrics_before_batch: bool,
498 ) {
499 let mut metrics = StreamMetrics::new(should_send_partial_metrics);
500 let mut last_metrics_str = None;
501 let recordbatch_schema = recordbatches.schema();
502 let schema = recordbatch_schema.arrow_schema().clone();
503 let start = Instant::now();
504 if let Err(e) = tx.send(Ok(FlightMessage::Schema(schema))).await {
505 warn!(e; "stop sending Flight data");
506 return;
507 }
508 metrics.send_schema_duration += start.elapsed();
509
510 let reached_eof = if should_send_partial_metrics {
514 Self::verbose_metrics_stream(
515 &mut recordbatches,
516 &mut tx,
517 &mut metrics,
518 &mut last_metrics_str,
519 can_send_metrics_before_batch,
520 )
521 .await
522 } else {
523 CoalescingBatcher::new(recordbatch_schema.clone())
524 .run(&mut recordbatches, &mut tx, &mut metrics)
525 .await
526 };
527 if !reached_eof {
528 return;
529 }
530
531 if let Some(metrics_str) = recordbatches
533 .metrics()
534 .and_then(|m| serde_json::to_string(&m).ok())
535 {
536 let _ = Self::send_metrics(&mut tx, &mut metrics, metrics_str).await;
537 }
538 }
539
540 async fn verbose_metrics_stream(
545 recordbatches: &mut SendableRecordBatchStream,
546 tx: &mut Sender<TonicResult<FlightMessage>>,
547 metrics: &mut StreamMetrics,
548 last_metrics_str: &mut Option<String>,
549 can_send_metrics_before_batch: bool,
550 ) -> bool {
551 loop {
552 let start = Instant::now();
553 let batch_or_err = if can_send_metrics_before_batch {
554 match time::timeout(
555 FLIGHT_METRICS_HEARTBEAT_INTERVAL,
556 recordbatches.next().in_current_span(),
557 )
558 .await
559 {
560 Ok(result) => result,
561 Err(_) => {
562 if let Some(metrics_str) = recordbatches
563 .metrics()
564 .and_then(|m| serde_json::to_string(&m).ok())
565 && !Self::send_metrics_if_changed(
566 tx,
567 metrics,
568 last_metrics_str,
569 metrics_str,
570 )
571 .await
572 {
573 return false;
574 }
575 metrics.fetch_content_duration += start.elapsed();
576 continue;
577 }
578 }
579 } else {
580 recordbatches.next().in_current_span().await
581 };
582 metrics.fetch_content_duration += start.elapsed();
583 let Some(batch_or_err) = batch_or_err else {
584 break;
585 };
586 match batch_or_err {
587 Ok(recordbatch) => {
588 metrics.total_rows += recordbatch.num_rows();
589 metrics.record_batch_count += 1;
590 metrics.total_bytes += recordbatch.df_record_batch().get_array_memory_size();
591 let start = Instant::now();
592 if let Err(e) = tx
593 .send(Ok(FlightMessage::RecordBatch(
594 recordbatch.into_df_record_batch(),
595 )))
596 .await
597 {
598 warn!(e; "stop sending Flight data");
599 return false;
600 }
601 metrics.send_record_batch_duration += start.elapsed();
602 if let Some(metrics_str) = recordbatches
603 .metrics()
604 .and_then(|m| serde_json::to_string(&m).ok())
605 && {
606 *last_metrics_str = Some(metrics_str.clone());
607 !Self::send_metrics(tx, metrics, metrics_str).await
608 }
609 {
610 return false;
611 }
612 }
613 Err(e) => {
614 if e.status_code().should_log_error() {
615 error!("{e:?}");
616 }
617 let e = Err(e).context(error::CollectRecordbatchSnafu);
618 if let Err(e) = tx.send(e.map_err(|x| x.into())).await {
619 warn!(e; "stop sending Flight data");
620 }
621 return false;
622 }
623 }
624 }
625 true
626 }
627}
628
629#[pinned_drop]
630impl PinnedDrop for FlightRecordBatchStream {
631 fn drop(self: Pin<&mut Self>) {
632 self.join_handle.abort();
633 }
634}
635
636impl Stream for FlightRecordBatchStream {
637 type Item = TonicResult<FlightData>;
638
639 fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
640 let this = self.project();
641 if *this.done {
642 Poll::Ready(None)
643 } else {
644 if let Some(x) = this.buffer.pop_front() {
645 return Poll::Ready(Some(Ok(x)));
646 }
647 match this.rx.poll_next(cx) {
648 Poll::Ready(None) => {
649 *this.done = true;
650 Poll::Ready(None)
651 }
652 Poll::Ready(Some(result)) => match result {
653 Ok(flight_message) => {
654 let mut iter = this.encoder.encode(flight_message).into_iter();
655 let Some(first) = iter.next() else {
656 unreachable!()
659 };
660 this.buffer.extend(iter);
661 Poll::Ready(Some(Ok(first)))
662 }
663 Err(e) => {
664 *this.done = true;
665 Poll::Ready(Some(Err(e)))
666 }
667 },
668 Poll::Pending => Poll::Pending,
669 }
670 }
671 }
672}
673
674#[cfg(test)]
675mod test {
676 use std::pin::Pin;
677 use std::sync::Arc;
678 use std::task::{Context, Poll};
679 use std::time::Duration;
680
681 use common_grpc::flight::{FlightDecoder, FlightMessage};
682 use common_recordbatch::adapter::RecordBatchMetrics;
683 use common_recordbatch::error::CreateRecordBatchesSnafu;
684 use common_recordbatch::{OrderOption, RecordBatch, RecordBatchStream, RecordBatches};
685 use datatypes::arrow::array::{ArrayRef, DictionaryArray, Int32Array, StringArray};
686 use datatypes::arrow::datatypes::Int32Type;
687 use datatypes::prelude::*;
688 use datatypes::schema::{ColumnSchema, Schema, SchemaRef};
689 use datatypes::vectors::{DictionaryVector, Int32Vector};
690 use futures::StreamExt;
691 use session::context::{
692 LIVE_ANALYZE_METRICS_EXTENSION_KEY, QueryContext,
693 SUPPORT_FLIGHT_METRICS_BEFORE_BATCH_EXTENSION_KEY,
694 };
695
696 use super::*;
697
698 struct PendingMetricsStream {
699 schema: SchemaRef,
700 metrics: RecordBatchMetrics,
701 }
702
703 struct MetricsThenBatchStream {
704 schema: SchemaRef,
705 metrics: RecordBatchMetrics,
706 rx: tokio::sync::mpsc::UnboundedReceiver<common_recordbatch::error::Result<RecordBatch>>,
707 }
708
709 enum ScriptedItem {
710 Batch(common_recordbatch::error::Result<RecordBatch>),
711 Pending,
712 PersistentPending,
713 }
714
715 struct ScriptedBatchStream {
716 schema: SchemaRef,
717 items: VecDeque<ScriptedItem>,
718 poll_count: Arc<std::sync::atomic::AtomicUsize>,
719 }
720
721 struct DropFlagStream {
722 schema: SchemaRef,
723 dropped: Arc<std::sync::atomic::AtomicBool>,
724 }
725
726 fn query_context_with_matching_capability() -> Arc<QueryContext> {
727 let query_ctx = QueryContext::arc();
728 let remote_query_id = query_ctx
729 .remote_query_id()
730 .expect("query context must have remote query id")
731 .to_string();
732 let mut query_ctx = (*query_ctx).clone();
733 query_ctx.set_extension(
734 SUPPORT_FLIGHT_METRICS_BEFORE_BATCH_EXTENSION_KEY,
735 remote_query_id,
736 );
737 Arc::new(query_ctx)
738 }
739
740 fn query_context_with_live_metrics_and_matching_capability() -> Arc<QueryContext> {
741 let mut query_ctx = (*query_context_with_matching_capability()).clone();
742 query_ctx.enable_live_analyze_metrics();
743 Arc::new(query_ctx)
744 }
745
746 impl RecordBatchStream for PendingMetricsStream {
747 fn schema(&self) -> SchemaRef {
748 self.schema.clone()
749 }
750
751 fn output_ordering(&self) -> Option<&[OrderOption]> {
752 None
753 }
754
755 fn metrics(&self) -> Option<RecordBatchMetrics> {
756 Some(self.metrics.clone())
757 }
758 }
759
760 impl Stream for PendingMetricsStream {
761 type Item = common_recordbatch::error::Result<RecordBatch>;
762
763 fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
764 Poll::Pending
765 }
766 }
767
768 impl RecordBatchStream for MetricsThenBatchStream {
769 fn schema(&self) -> SchemaRef {
770 self.schema.clone()
771 }
772
773 fn output_ordering(&self) -> Option<&[OrderOption]> {
774 None
775 }
776
777 fn metrics(&self) -> Option<RecordBatchMetrics> {
778 Some(self.metrics.clone())
779 }
780 }
781
782 impl Stream for MetricsThenBatchStream {
783 type Item = common_recordbatch::error::Result<RecordBatch>;
784
785 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
786 self.rx.poll_recv(cx)
787 }
788 }
789
790 impl RecordBatchStream for ScriptedBatchStream {
791 fn schema(&self) -> SchemaRef {
792 self.schema.clone()
793 }
794
795 fn output_ordering(&self) -> Option<&[OrderOption]> {
796 None
797 }
798
799 fn metrics(&self) -> Option<RecordBatchMetrics> {
800 Some(RecordBatchMetrics::default())
801 }
802 }
803
804 impl Stream for ScriptedBatchStream {
805 type Item = common_recordbatch::error::Result<RecordBatch>;
806
807 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
808 self.poll_count
809 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
810 match self.items.pop_front() {
811 Some(ScriptedItem::Batch(item)) => Poll::Ready(Some(item)),
812 Some(ScriptedItem::Pending) => {
813 cx.waker().wake_by_ref();
814 Poll::Pending
815 }
816 Some(ScriptedItem::PersistentPending) => {
817 self.items.push_front(ScriptedItem::PersistentPending);
818 Poll::Pending
819 }
820 None => Poll::Ready(None),
821 }
822 }
823 }
824
825 impl RecordBatchStream for DropFlagStream {
826 fn schema(&self) -> SchemaRef {
827 self.schema.clone()
828 }
829
830 fn output_ordering(&self) -> Option<&[OrderOption]> {
831 None
832 }
833
834 fn metrics(&self) -> Option<RecordBatchMetrics> {
835 None
836 }
837 }
838
839 impl Stream for DropFlagStream {
840 type Item = common_recordbatch::error::Result<RecordBatch>;
841
842 fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
843 Poll::Pending
844 }
845 }
846
847 impl Drop for DropFlagStream {
848 fn drop(&mut self) {
849 self.dropped
850 .store(true, std::sync::atomic::Ordering::Relaxed);
851 }
852 }
853
854 fn int_batch(schema: SchemaRef, values: impl IntoIterator<Item = i32>) -> RecordBatch {
855 RecordBatch::new(
856 schema,
857 vec![Arc::new(Int32Vector::from_iter_values(values)) as VectorRef],
858 )
859 .unwrap()
860 }
861
862 async fn flight_messages_with_context(
863 recordbatches: SendableRecordBatchStream,
864 query_ctx: Arc<QueryContext>,
865 ) -> Vec<TonicResult<FlightMessage>> {
866 let mut stream = FlightRecordBatchStream::new(
867 FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
868 recordbatches,
869 )),
870 TracingContext::default(),
871 FlightCompression::default(),
872 query_ctx,
873 );
874 let decoder = &mut FlightDecoder::default();
875 let mut messages = Vec::new();
876 while let Some(data) = stream.next().await {
877 match data {
878 Ok(data) => {
879 if let Some(message) = decoder.try_decode(&data).unwrap() {
880 messages.push(Ok(message));
881 }
882 }
883 Err(status) => messages.push(Err(status)),
884 }
885 }
886 messages
887 }
888
889 async fn flight_messages(
890 recordbatches: SendableRecordBatchStream,
891 ) -> Vec<TonicResult<FlightMessage>> {
892 flight_messages_with_context(recordbatches, QueryContext::arc()).await
893 }
894
895 #[tokio::test]
896 async fn test_drop_cancels_and_releases_upstream_stream() {
897 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
898 "a",
899 ConcreteDataType::int32_datatype(),
900 false,
901 )]));
902 let dropped = Arc::new(std::sync::atomic::AtomicBool::new(false));
903 let stream = FlightRecordBatchStream::new(
904 FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(Box::pin(
905 DropFlagStream {
906 schema,
907 dropped: dropped.clone(),
908 },
909 ))),
910 TracingContext::default(),
911 FlightCompression::default(),
912 QueryContext::arc(),
913 );
914 drop(stream);
915 tokio::time::timeout(Duration::from_secs(1), async {
916 while !dropped.load(std::sync::atomic::Ordering::Relaxed) {
917 tokio::task::yield_now().await;
918 }
919 })
920 .await
921 .expect("dropping Flight stream must release upstream");
922 }
923
924 #[tokio::test]
925 async fn test_first_batch_is_sent_before_any_additional_poll() {
926 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
927 "a",
928 ConcreteDataType::int32_datatype(),
929 false,
930 )]));
931 let poll_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
932 let recordbatches: SendableRecordBatchStream = Box::pin(ScriptedBatchStream {
933 schema: schema.clone(),
934 items: VecDeque::from([
935 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
936 ScriptedItem::PersistentPending,
937 ]),
938 poll_count: poll_count.clone(),
939 });
940 let (tx, mut rx) = mpsc::channel::<TonicResult<FlightMessage>>(1);
947 let handle = tokio::spawn(FlightRecordBatchStream::flight_data_stream(
948 recordbatches,
949 tx,
950 false,
951 false,
952 ));
953 tokio::time::timeout(Duration::from_secs(1), async {
954 while poll_count.load(std::sync::atomic::Ordering::Relaxed) < 1 {
955 tokio::task::yield_now().await;
956 }
957 })
958 .await
959 .expect("the producer must fetch the first batch");
960 assert_eq!(
964 poll_count.load(std::sync::atomic::Ordering::Relaxed),
965 1,
966 "the first batch must be sent before polling the source again"
967 );
968 assert!(matches!(
969 rx.next().await.unwrap().unwrap(),
970 FlightMessage::Schema(_)
971 ));
972 let first = rx.next().await.unwrap().unwrap();
973 let FlightMessage::RecordBatch(batch) = first else {
974 panic!("expected the first record batch");
975 };
976 assert_eq!(batch.num_rows(), 1);
977 handle.abort();
978 }
979
980 #[tokio::test]
981 async fn test_ready_only_coalesces_later_batches_without_polling_after_first() {
982 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
983 "a",
984 ConcreteDataType::int32_datatype(),
985 false,
986 )]));
987 let poll_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
988 let recordbatches = Box::pin(ScriptedBatchStream {
989 schema: schema.clone(),
990 items: VecDeque::from([
991 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
992 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [2]))),
993 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [3]))),
994 ]),
995 poll_count: poll_count.clone(),
996 });
997
998 let messages = flight_messages(recordbatches).await;
999 assert_eq!(poll_count.load(std::sync::atomic::Ordering::Relaxed), 4);
1000 assert!(matches!(messages[0], Ok(FlightMessage::Schema(_))));
1001 let FlightMessage::RecordBatch(first) = messages[1].as_ref().unwrap() else {
1002 panic!("expected the first record batch");
1003 };
1004 assert_eq!(first.num_rows(), 1);
1005 let FlightMessage::RecordBatch(merged) = messages[2].as_ref().unwrap() else {
1006 panic!("expected the coalesced record batch");
1007 };
1008 assert_eq!(merged.num_rows(), 2);
1009 }
1010
1011 #[tokio::test]
1022 async fn test_merged_source_batches_are_released_before_send() {
1023 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1024 "a",
1025 ConcreteDataType::int32_datatype(),
1026 false,
1027 )]));
1028 let tracked: ArrayRef = Arc::new(Int32Array::from(vec![2]));
1030 let weak = Arc::downgrade(&tracked);
1031 let tracked_batch = RecordBatch::from_df_record_batch(
1032 schema.clone(),
1033 common_recordbatch::DfRecordBatch::try_new(
1034 schema.arrow_schema().clone(),
1035 vec![tracked.clone()],
1036 )
1037 .unwrap(),
1038 );
1039 drop(tracked);
1042
1043 let mut items = vec![ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1])))];
1046 items.push(ScriptedItem::Batch(Ok(tracked_batch)));
1047 items.extend((0..15).map(|_| ScriptedItem::Batch(Ok(int_batch(schema.clone(), [3])))));
1048 let poll_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
1049 let recordbatches: SendableRecordBatchStream = Box::pin(ScriptedBatchStream {
1050 schema: schema.clone(),
1051 items: items.into(),
1052 poll_count: poll_count.clone(),
1053 });
1054 let (tx, mut rx) = mpsc::channel::<TonicResult<FlightMessage>>(1);
1055 let handle = tokio::spawn(FlightRecordBatchStream::flight_data_stream(
1056 recordbatches,
1057 tx,
1058 false,
1059 false,
1060 ));
1061 assert!(matches!(
1062 rx.next().await.unwrap().unwrap(),
1063 FlightMessage::Schema(_)
1064 ));
1065
1066 tokio::time::timeout(Duration::from_secs(1), async {
1068 while poll_count.load(std::sync::atomic::Ordering::Relaxed) != 17 {
1069 tokio::task::yield_now().await;
1070 }
1071 })
1072 .await
1073 .expect("the 16-batch group must be fetched");
1074 tokio::time::timeout(Duration::from_secs(1), async {
1077 while weak.upgrade().is_some() {
1078 tokio::task::yield_now().await;
1079 }
1080 })
1081 .await
1082 .expect("the merged source batches must be released before the merged group is sent");
1083 assert_eq!(
1084 poll_count.load(std::sync::atomic::Ordering::Relaxed),
1085 17,
1086 "the producer must still be blocked on the merged send"
1087 );
1088 handle.abort();
1089 }
1090
1091 #[tokio::test]
1092 async fn test_ready_only_exact_row_cap_flushes() {
1093 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1094 "a",
1095 ConcreteDataType::int32_datatype(),
1096 false,
1097 )]));
1098 let poll_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
1099 let recordbatches = Box::pin(ScriptedBatchStream {
1100 schema: schema.clone(),
1101 items: VecDeque::from([
1102 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
1103 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
1104 ScriptedItem::Batch(Ok(int_batch(schema.clone(), 0..4095))),
1105 ScriptedItem::PersistentPending,
1106 ]),
1107 poll_count: poll_count.clone(),
1108 });
1109 let (tx, mut rx) = mpsc::channel::<TonicResult<FlightMessage>>(1);
1110 let handle = tokio::spawn(FlightRecordBatchStream::flight_data_stream(
1111 recordbatches,
1112 tx,
1113 false,
1114 false,
1115 ));
1116
1117 assert!(matches!(
1118 rx.next().await.unwrap().unwrap(),
1119 FlightMessage::Schema(_)
1120 ));
1121 tokio::time::timeout(Duration::from_secs(1), async {
1122 while poll_count.load(std::sync::atomic::Ordering::Relaxed) != 3 {
1123 tokio::task::yield_now().await;
1124 }
1125 })
1126 .await
1127 .expect("exact-cap group must be formed");
1128 assert!(matches!(
1129 rx.next().await.unwrap().unwrap(),
1130 FlightMessage::RecordBatch(_)
1131 ));
1132 let FlightMessage::RecordBatch(batch) = rx.next().await.unwrap().unwrap() else {
1133 panic!("expected the exact-cap group");
1134 };
1135 assert_eq!(batch.num_rows(), 4096);
1136 handle.abort();
1137 }
1138
1139 #[tokio::test]
1140 async fn test_over_cap_batch_inside_group_is_sent_as_singleton() {
1141 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1142 "a",
1143 ConcreteDataType::int32_datatype(),
1144 false,
1145 )]));
1146 let poll_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
1147 let messages = flight_messages(Box::pin(ScriptedBatchStream {
1148 schema: schema.clone(),
1149 items: VecDeque::from([
1150 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
1151 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [2]))),
1152 ScriptedItem::Batch(Ok(int_batch(schema.clone(), 0..4096))),
1153 ]),
1154 poll_count: poll_count.clone(),
1155 }))
1156 .await;
1157 let batches = messages
1158 .iter()
1159 .filter_map(|message| match message.as_ref().unwrap() {
1160 FlightMessage::RecordBatch(batch) => Some(batch.num_rows()),
1161 _ => None,
1162 })
1163 .collect::<Vec<_>>();
1164 assert_eq!(batches, vec![1, 1, 4096]);
1169 assert_eq!(poll_count.load(std::sync::atomic::Ordering::Relaxed), 4);
1171 }
1172
1173 #[tokio::test]
1177 async fn test_tiny_group_then_over_budget_batch_are_sent_separately() {
1178 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1179 "a",
1180 ConcreteDataType::int32_datatype(),
1181 false,
1182 )]));
1183 let messages = flight_messages(Box::pin(ScriptedBatchStream {
1184 schema: schema.clone(),
1185 items: VecDeque::from([
1186 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
1187 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [2]))),
1188 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [3]))),
1189 ScriptedItem::Batch(Ok(int_batch(schema.clone(), 0..4096))),
1190 ]),
1191 poll_count: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
1192 }))
1193 .await;
1194 let batches = messages
1195 .iter()
1196 .filter_map(|message| match message.as_ref().unwrap() {
1197 FlightMessage::RecordBatch(batch) => Some(batch.num_rows()),
1198 _ => None,
1199 })
1200 .collect::<Vec<_>>();
1201 assert_eq!(batches, vec![1, 2, 4096]);
1205 }
1206
1207 #[tokio::test]
1208 async fn test_ready_only_soft_row_budget_flushes_oversized_group() {
1209 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1210 "a",
1211 ConcreteDataType::int32_datatype(),
1212 false,
1213 )]));
1214 let poll_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
1215 let messages = flight_messages(Box::pin(ScriptedBatchStream {
1216 schema: schema.clone(),
1217 items: VecDeque::from([
1218 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
1219 ScriptedItem::Batch(Ok(int_batch(schema.clone(), 0..2100))),
1220 ScriptedItem::Batch(Ok(int_batch(schema.clone(), 0..2100))),
1221 ]),
1222 poll_count: poll_count.clone(),
1223 }))
1224 .await;
1225 let batches = messages
1226 .iter()
1227 .filter_map(|message| match message.as_ref().unwrap() {
1228 FlightMessage::RecordBatch(batch) => Some(batch.num_rows()),
1229 _ => None,
1230 })
1231 .collect::<Vec<_>>();
1232 assert_eq!(batches, vec![1, 4200]);
1235 assert_eq!(poll_count.load(std::sync::atomic::Ordering::Relaxed), 4);
1237 }
1238
1239 #[tokio::test]
1240 async fn test_ready_only_current_at_cap_batch_is_sent_as_singleton() {
1241 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1242 "a",
1243 ConcreteDataType::int32_datatype(),
1244 false,
1245 )]));
1246 let poll_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
1247 let messages = flight_messages(Box::pin(ScriptedBatchStream {
1248 schema: schema.clone(),
1249 items: VecDeque::from([
1250 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
1251 ScriptedItem::Batch(Ok(int_batch(schema.clone(), 0..4096))),
1252 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [2]))),
1253 ]),
1254 poll_count: poll_count.clone(),
1255 }))
1256 .await;
1257 let batches = messages
1258 .iter()
1259 .filter_map(|message| match message.as_ref().unwrap() {
1260 FlightMessage::RecordBatch(batch) => Some(batch.num_rows()),
1261 _ => None,
1262 })
1263 .collect::<Vec<_>>();
1264 assert_eq!(batches, vec![1, 4096, 1]);
1267 assert_eq!(poll_count.load(std::sync::atomic::Ordering::Relaxed), 4);
1268 }
1269
1270 #[tokio::test]
1271 async fn test_ready_only_flushes_before_pending_and_before_error() {
1272 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1273 "a",
1274 ConcreteDataType::int32_datatype(),
1275 false,
1276 )]));
1277 let poll_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
1278 let recordbatches = Box::pin(ScriptedBatchStream {
1279 schema: schema.clone(),
1280 items: VecDeque::from([
1281 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
1282 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [2]))),
1283 ScriptedItem::Pending,
1284 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [3]))),
1285 ScriptedItem::Batch(Err(CreateRecordBatchesSnafu {
1286 reason: "expected failure".to_string(),
1287 }
1288 .build())),
1289 ]),
1290 poll_count,
1291 });
1292
1293 let messages = flight_messages(recordbatches).await;
1294 assert!(matches!(messages[1], Ok(FlightMessage::RecordBatch(_))));
1295 assert!(matches!(messages[2], Ok(FlightMessage::RecordBatch(_))));
1296 assert!(matches!(messages[3], Ok(FlightMessage::RecordBatch(_))));
1297 assert!(messages[4].is_err());
1298 assert_eq!(messages.len(), 5);
1299 }
1300
1301 struct MetricsCountingErrorStream {
1307 schema: SchemaRef,
1308 yielded_batch: bool,
1309 metrics_calls: Arc<std::sync::atomic::AtomicUsize>,
1310 }
1311
1312 impl RecordBatchStream for MetricsCountingErrorStream {
1313 fn schema(&self) -> SchemaRef {
1314 self.schema.clone()
1315 }
1316
1317 fn output_ordering(&self) -> Option<&[OrderOption]> {
1318 None
1319 }
1320
1321 fn metrics(&self) -> Option<RecordBatchMetrics> {
1322 self.metrics_calls
1323 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1324 Some(RecordBatchMetrics::default())
1325 }
1326 }
1327
1328 impl Stream for MetricsCountingErrorStream {
1329 type Item = common_recordbatch::error::Result<RecordBatch>;
1330
1331 fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1332 if self.yielded_batch {
1333 return Poll::Ready(Some(Err(CreateRecordBatchesSnafu {
1334 reason: "expected failure".to_string(),
1335 }
1336 .build())));
1337 }
1338 self.yielded_batch = true;
1339 Poll::Ready(Some(Ok(int_batch(self.schema.clone(), [1]))))
1340 }
1341 }
1342
1343 #[tokio::test]
1349 async fn test_verbose_error_skips_final_metrics_tail() {
1350 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1351 "a",
1352 ConcreteDataType::int32_datatype(),
1353 false,
1354 )]));
1355 let metrics_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
1356 let recordbatches: SendableRecordBatchStream = Box::pin(MetricsCountingErrorStream {
1357 schema: schema.clone(),
1358 yielded_batch: false,
1359 metrics_calls: metrics_calls.clone(),
1360 });
1361 let (tx, mut rx) = mpsc::channel::<TonicResult<FlightMessage>>(8);
1362 FlightRecordBatchStream::flight_data_stream(recordbatches, tx, true, false).await;
1365 assert_eq!(metrics_calls.load(std::sync::atomic::Ordering::Relaxed), 1);
1369 let mut messages = Vec::new();
1371 while let Some(msg) = rx.next().await {
1372 messages.push(msg);
1373 }
1374 assert!(matches!(messages[0], Ok(FlightMessage::Schema(_))));
1375 assert!(matches!(messages[1], Ok(FlightMessage::RecordBatch(_))));
1376 assert!(matches!(messages[2], Ok(FlightMessage::Metrics(_))));
1377 assert!(messages[3].is_err());
1378 assert_eq!(messages.len(), 4);
1379 }
1380
1381 #[tokio::test]
1382 async fn test_verbose_ready_batches_preserve_batch_metrics_order() {
1383 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1384 "a",
1385 ConcreteDataType::int32_datatype(),
1386 false,
1387 )]));
1388 let query_ctx = QueryContext::arc();
1389 query_ctx.set_explain_verbose(true);
1390 let messages = flight_messages_with_context(
1391 Box::pin(ScriptedBatchStream {
1392 schema: schema.clone(),
1393 items: VecDeque::from([
1394 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
1395 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [2]))),
1396 ]),
1397 poll_count: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
1398 }),
1399 query_ctx,
1400 )
1401 .await;
1402 assert!(matches!(messages[1], Ok(FlightMessage::RecordBatch(_))));
1403 assert!(matches!(messages[2], Ok(FlightMessage::Metrics(_))));
1404 assert!(matches!(messages[3], Ok(FlightMessage::RecordBatch(_))));
1405 assert!(matches!(messages[4], Ok(FlightMessage::Metrics(_))));
1406 }
1407
1408 #[tokio::test]
1409 async fn test_ready_only_flushes_group_before_at_cap_batch_and_coalesces_empty_batches() {
1410 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1411 "a",
1412 ConcreteDataType::int32_datatype(),
1413 false,
1414 )]));
1415 let oversized = int_batch(schema.clone(), 0..4096);
1416 let mut items = vec![
1417 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [1]))),
1418 ScriptedItem::Batch(Ok(int_batch(schema.clone(), [2]))),
1419 ScriptedItem::Batch(Ok(oversized)),
1420 ];
1421 items.extend(
1425 (0..17).map(|_| ScriptedItem::Batch(Ok(RecordBatch::new_empty(schema.clone())))),
1426 );
1427 let poll_count = Arc::new(std::sync::atomic::AtomicUsize::new(0));
1428 let messages = flight_messages(Box::pin(ScriptedBatchStream {
1429 schema,
1430 items: items.into(),
1431 poll_count: poll_count.clone(),
1432 }))
1433 .await;
1434 let batches = messages
1435 .iter()
1436 .filter_map(|message| match message.as_ref().unwrap() {
1437 FlightMessage::RecordBatch(batch) => Some(batch.num_rows()),
1438 _ => None,
1439 })
1440 .collect::<Vec<_>>();
1441 assert_eq!(batches, vec![1, 1, 4096, 0, 0]);
1445 assert_eq!(poll_count.load(std::sync::atomic::Ordering::Relaxed), 21);
1446 }
1447
1448 #[tokio::test]
1449 async fn test_ready_only_flushes_group_before_byte_oversized_batch() {
1450 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1451 "a",
1452 ConcreteDataType::string_datatype(),
1453 false,
1454 )]));
1455 let oversized = RecordBatch::new(
1456 schema.clone(),
1457 vec![Arc::new(datatypes::vectors::StringVector::from_slice(
1458 &vec!["x".repeat(300); 1023],
1459 )) as VectorRef],
1460 )
1461 .unwrap();
1462 let messages = flight_messages(Box::pin(ScriptedBatchStream {
1463 schema: schema.clone(),
1464 items: VecDeque::from([
1465 ScriptedItem::Batch(Ok(RecordBatch::new_empty(schema.clone()))),
1466 ScriptedItem::Batch(Ok(RecordBatch::new_empty(schema.clone()))),
1467 ScriptedItem::Batch(Ok(oversized)),
1468 ]),
1469 poll_count: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
1470 }))
1471 .await;
1472 let batches = messages
1473 .iter()
1474 .filter_map(|message| match message.as_ref().unwrap() {
1475 FlightMessage::RecordBatch(batch) => Some(batch.num_rows()),
1476 _ => None,
1477 })
1478 .collect::<Vec<_>>();
1479 assert_eq!(batches, vec![0, 0, 1023]);
1483 }
1484
1485 #[tokio::test]
1486 async fn test_ready_only_coalesces_dictionary_batches_with_different_mappings() {
1487 let arrow_schema = Arc::new(datatypes::arrow::datatypes::Schema::new(vec![
1488 datatypes::arrow::datatypes::Field::new_dictionary(
1489 "a",
1490 datatypes::arrow::datatypes::DataType::Int32,
1491 datatypes::arrow::datatypes::DataType::Utf8,
1492 false,
1493 ),
1494 ]));
1495 let schema = Arc::new(Schema::try_from(arrow_schema).unwrap());
1496 let dictionary_batch = |keys, values| {
1497 let array = DictionaryArray::<Int32Type>::new(
1498 Int32Array::from(keys),
1499 Arc::new(StringArray::from(values)),
1500 );
1501 RecordBatch::new(
1502 schema.clone(),
1503 vec![Arc::new(
1504 DictionaryVector::new(array, ConcreteDataType::string_datatype()).unwrap(),
1505 ) as VectorRef],
1506 )
1507 .unwrap()
1508 };
1509 let recordbatches = Box::pin(ScriptedBatchStream {
1510 schema: schema.clone(),
1511 items: VecDeque::from([
1512 ScriptedItem::Batch(Ok(dictionary_batch(vec![0], vec!["zero"]))),
1513 ScriptedItem::Batch(Ok(dictionary_batch(vec![0], vec!["first"]))),
1514 ScriptedItem::Batch(Ok(dictionary_batch(vec![0], vec!["second"]))),
1515 ]),
1516 poll_count: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
1517 });
1518 let messages = flight_messages(recordbatches).await;
1519 let FlightMessage::RecordBatch(merged) = messages[2].as_ref().unwrap() else {
1520 panic!("expected the coalesced dictionary batch");
1521 };
1522 assert_eq!(merged.num_rows(), 2);
1523 let dictionary = merged
1530 .column(0)
1531 .as_any()
1532 .downcast_ref::<DictionaryArray<Int32Type>>()
1533 .expect("expected a dictionary column");
1534 let values = dictionary
1535 .values()
1536 .as_any()
1537 .downcast_ref::<StringArray>()
1538 .expect("expected a dictionary of Utf8 values");
1539 let logical = (0..merged.num_rows())
1540 .map(|row| values.value(dictionary.keys().value(row) as usize))
1541 .collect::<Vec<_>>();
1542 assert_eq!(logical, vec!["first", "second"]);
1543 }
1544
1545 #[tokio::test]
1546 async fn test_flight_record_batch_stream() {
1547 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1548 "a",
1549 ConcreteDataType::int32_datatype(),
1550 false,
1551 )]));
1552
1553 let v: VectorRef = Arc::new(Int32Vector::from_slice([1, 2]));
1554 let recordbatch = RecordBatch::new(schema.clone(), vec![v]).unwrap();
1555
1556 let recordbatches = RecordBatches::try_new(schema.clone(), vec![recordbatch.clone()])
1557 .unwrap()
1558 .as_stream();
1559 let mut stream = FlightRecordBatchStream::new(
1560 FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
1561 recordbatches,
1562 )),
1563 TracingContext::default(),
1564 FlightCompression::default(),
1565 QueryContext::arc(),
1566 );
1567
1568 let mut raw_data = Vec::with_capacity(2);
1569 raw_data.push(stream.next().await.unwrap().unwrap());
1570 raw_data.push(stream.next().await.unwrap().unwrap());
1571 assert!(stream.next().await.is_none());
1572 assert!(stream.done);
1573
1574 let decoder = &mut FlightDecoder::default();
1575 let mut flight_messages = raw_data
1576 .into_iter()
1577 .map(|x| decoder.try_decode(&x).unwrap().unwrap())
1578 .collect::<Vec<FlightMessage>>();
1579 assert_eq!(flight_messages.len(), 2);
1580
1581 match flight_messages.remove(0) {
1582 FlightMessage::Schema(actual_schema) => {
1583 assert_eq!(&actual_schema, schema.arrow_schema());
1584 }
1585 _ => unreachable!(),
1586 }
1587
1588 match flight_messages.remove(0) {
1589 FlightMessage::RecordBatch(actual_recordbatch) => {
1590 assert_eq!(&actual_recordbatch, recordbatch.df_record_batch());
1591 }
1592 _ => unreachable!(),
1593 }
1594 }
1595
1596 #[tokio::test]
1597 async fn test_flight_record_batch_stream_encodes_affected_rows() {
1598 let mut stream = FlightRecordBatchStream::new(
1599 FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::AffectedRows {
1600 rows: 42,
1601 metrics: Some(r#"{"region_watermarks":[]}"#.to_string()),
1602 }),
1603 TracingContext::default(),
1604 FlightCompression::default(),
1605 QueryContext::arc(),
1606 );
1607
1608 let data = stream.next().await.unwrap().unwrap();
1609 let message = FlightDecoder::default().try_decode(&data).unwrap().unwrap();
1610 assert!(matches!(
1611 message,
1612 FlightMessage::AffectedRows {
1613 rows: 42,
1614 metrics: Some(_),
1615 }
1616 ));
1617 assert!(stream.next().await.is_none());
1618 }
1619
1620 #[tokio::test]
1621 async fn test_flight_record_batch_stream_forwards_initializer_error() {
1622 let mut stream = FlightRecordBatchStream::new(
1623 FlightRecordBatchStreamInput::initializer(async {
1624 Err(tonic::Status::unavailable(
1625 "remote read initialization failed",
1626 ))
1627 }),
1628 TracingContext::default(),
1629 FlightCompression::default(),
1630 QueryContext::arc(),
1631 );
1632
1633 let error = stream.next().await.unwrap().unwrap_err();
1634 assert_eq!(tonic::Code::Unavailable, error.code());
1635 assert!(stream.next().await.is_none());
1636 }
1637 #[tokio::test]
1638 async fn test_flight_record_batch_stream_emits_metrics_while_pending() {
1639 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1640 "a",
1641 ConcreteDataType::int32_datatype(),
1642 false,
1643 )]));
1644 let metrics = RecordBatchMetrics {
1645 elapsed_compute: 42,
1646 ..Default::default()
1647 };
1648 let recordbatches = Box::pin(PendingMetricsStream {
1649 schema: schema.clone(),
1650 metrics,
1651 });
1652 let query_ctx = query_context_with_live_metrics_and_matching_capability();
1653 let initializer_query_ctx = query_ctx.clone();
1654 let mut stream = FlightRecordBatchStream::new(
1655 FlightRecordBatchStreamInput::initializer(async move {
1656 initializer_query_ctx.set_explain_verbose(true);
1657 Ok(FlightRecordBatchSource::RecordBatches(recordbatches))
1658 }),
1659 TracingContext::default(),
1660 FlightCompression::default(),
1661 query_ctx,
1662 );
1663
1664 let decoder = &mut FlightDecoder::default();
1665 let schema_data = stream.next().await.unwrap().unwrap();
1666 match decoder.try_decode(&schema_data).unwrap().unwrap() {
1667 FlightMessage::Schema(actual_schema) => {
1668 assert_eq!(&actual_schema, schema.arrow_schema());
1669 }
1670 _ => unreachable!(),
1671 }
1672
1673 let metrics_data = tokio::time::timeout(Duration::from_secs(2), stream.next())
1674 .await
1675 .unwrap()
1676 .unwrap()
1677 .unwrap();
1678 match decoder.try_decode(&metrics_data).unwrap().unwrap() {
1679 FlightMessage::Metrics(metrics) => {
1680 let metrics: RecordBatchMetrics = serde_json::from_str(&metrics).unwrap();
1681 assert_eq!(metrics.elapsed_compute, 42);
1682 }
1683 other => panic!("expected metrics message, got {other:?}"),
1684 }
1685 }
1686
1687 #[tokio::test]
1688 async fn test_flight_record_batch_stream_continues_after_pending_metrics() {
1689 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1690 "a",
1691 ConcreteDataType::int32_datatype(),
1692 false,
1693 )]));
1694 let metrics = RecordBatchMetrics {
1695 elapsed_compute: 42,
1696 ..Default::default()
1697 };
1698 let recordbatch = RecordBatch::new(
1699 schema.clone(),
1700 vec![Arc::new(Int32Vector::from_slice([1])) as VectorRef],
1701 )
1702 .unwrap();
1703 let expected_recordbatch = recordbatch.df_record_batch().clone();
1704 let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
1705 let recordbatches = Box::pin(MetricsThenBatchStream {
1706 schema: schema.clone(),
1707 metrics,
1708 rx,
1709 });
1710 let query_ctx = query_context_with_live_metrics_and_matching_capability();
1711 query_ctx.set_explain_verbose(true);
1712 let mut stream = FlightRecordBatchStream::new(
1713 FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
1714 recordbatches,
1715 )),
1716 TracingContext::default(),
1717 FlightCompression::default(),
1718 query_ctx,
1719 );
1720
1721 let decoder = &mut FlightDecoder::default();
1722 let schema_data = stream.next().await.unwrap().unwrap();
1723 assert!(matches!(
1724 decoder.try_decode(&schema_data).unwrap().unwrap(),
1725 FlightMessage::Schema(_)
1726 ));
1727
1728 let metrics_data = tokio::time::timeout(Duration::from_secs(2), stream.next())
1729 .await
1730 .unwrap()
1731 .unwrap()
1732 .unwrap();
1733 assert!(matches!(
1734 decoder.try_decode(&metrics_data).unwrap().unwrap(),
1735 FlightMessage::Metrics(_)
1736 ));
1737
1738 tx.send(Ok(recordbatch)).unwrap();
1739 let batch_data = tokio::time::timeout(Duration::from_secs(2), stream.next())
1740 .await
1741 .unwrap()
1742 .unwrap()
1743 .unwrap();
1744 match decoder.try_decode(&batch_data).unwrap().unwrap() {
1745 FlightMessage::RecordBatch(actual_recordbatch) => {
1746 assert_eq!(&actual_recordbatch, &expected_recordbatch);
1747 }
1748 other => panic!("expected record batch after pending metrics, got {other:?}"),
1749 }
1750
1751 drop(tx);
1752 let final_metrics_data = tokio::time::timeout(Duration::from_secs(2), stream.next())
1753 .await
1754 .unwrap()
1755 .unwrap()
1756 .unwrap();
1757 assert!(matches!(
1758 decoder.try_decode(&final_metrics_data).unwrap().unwrap(),
1759 FlightMessage::Metrics(_)
1760 ));
1761 }
1762
1763 #[tokio::test]
1764 async fn test_flight_record_batch_stream_requires_live_metrics_for_pre_batch_metrics() {
1765 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1766 "a",
1767 ConcreteDataType::int32_datatype(),
1768 false,
1769 )]));
1770 let recordbatches = Box::pin(PendingMetricsStream {
1771 schema: schema.clone(),
1772 metrics: RecordBatchMetrics {
1773 elapsed_compute: 42,
1774 ..Default::default()
1775 },
1776 });
1777 let query_ctx = query_context_with_matching_capability();
1778 query_ctx.set_explain_verbose(true);
1779 let mut stream = FlightRecordBatchStream::new(
1780 FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
1781 recordbatches,
1782 )),
1783 TracingContext::default(),
1784 FlightCompression::default(),
1785 query_ctx,
1786 );
1787
1788 let decoder = &mut FlightDecoder::default();
1789 let schema_data = stream.next().await.unwrap().unwrap();
1790 assert!(matches!(
1791 decoder.try_decode(&schema_data).unwrap().unwrap(),
1792 FlightMessage::Schema(_)
1793 ));
1794 assert!(
1795 tokio::time::timeout(
1796 FLIGHT_METRICS_HEARTBEAT_INTERVAL + Duration::from_millis(200),
1797 stream.next()
1798 )
1799 .await
1800 .is_err(),
1801 "pre-batch Metrics must be gated by live analyze metrics"
1802 );
1803 }
1804
1805 #[tokio::test]
1806 async fn test_flight_record_batch_stream_rejects_spoofed_live_metrics_for_pre_batch_metrics() {
1807 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1808 "a",
1809 ConcreteDataType::int32_datatype(),
1810 false,
1811 )]));
1812 let recordbatches = Box::pin(PendingMetricsStream {
1813 schema: schema.clone(),
1814 metrics: RecordBatchMetrics {
1815 elapsed_compute: 42,
1816 ..Default::default()
1817 },
1818 });
1819 let query_ctx = query_context_with_live_metrics_and_matching_capability();
1820 let mut query_ctx = (*query_ctx).clone();
1821 query_ctx.set_extension(LIVE_ANALYZE_METRICS_EXTENSION_KEY, "true");
1822 let query_ctx = Arc::new(query_ctx);
1823 query_ctx.set_explain_verbose(true);
1824 let mut stream = FlightRecordBatchStream::new(
1825 FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
1826 recordbatches,
1827 )),
1828 TracingContext::default(),
1829 FlightCompression::default(),
1830 query_ctx,
1831 );
1832
1833 let decoder = &mut FlightDecoder::default();
1834 let schema_data = stream.next().await.unwrap().unwrap();
1835 assert!(matches!(
1836 decoder.try_decode(&schema_data).unwrap().unwrap(),
1837 FlightMessage::Schema(_)
1838 ));
1839 assert!(
1840 tokio::time::timeout(
1841 FLIGHT_METRICS_HEARTBEAT_INTERVAL + Duration::from_millis(200),
1842 stream.next()
1843 )
1844 .await
1845 .is_err(),
1846 "pre-batch Metrics must reject spoofed live analyze metrics"
1847 );
1848 }
1849
1850 #[tokio::test]
1851 async fn test_flight_record_batch_stream_requires_explain_verbose_for_pre_batch_metrics() {
1852 let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1853 "a",
1854 ConcreteDataType::int32_datatype(),
1855 false,
1856 )]));
1857 let recordbatches = Box::pin(PendingMetricsStream {
1858 schema: schema.clone(),
1859 metrics: RecordBatchMetrics {
1860 elapsed_compute: 42,
1861 ..Default::default()
1862 },
1863 });
1864 let query_ctx = query_context_with_matching_capability();
1865 let mut stream = FlightRecordBatchStream::new(
1866 FlightRecordBatchStreamInput::ready(FlightRecordBatchSource::RecordBatches(
1867 recordbatches,
1868 )),
1869 TracingContext::default(),
1870 FlightCompression::default(),
1871 query_ctx,
1872 );
1873
1874 let decoder = &mut FlightDecoder::default();
1875 let schema_data = stream.next().await.unwrap().unwrap();
1876 assert!(matches!(
1877 decoder.try_decode(&schema_data).unwrap().unwrap(),
1878 FlightMessage::Schema(_)
1879 ));
1880 assert!(
1881 tokio::time::timeout(
1882 FLIGHT_METRICS_HEARTBEAT_INTERVAL + Duration::from_millis(200),
1883 stream.next()
1884 )
1885 .await
1886 .is_err(),
1887 "pre-batch Metrics must be gated by explain verbose even when capability is set"
1888 );
1889 }
1890}