Skip to main content

servers/grpc/flight/
stream.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::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
55/// Determines whether a Flight result is ready now or initialized asynchronously.
56pub enum FlightRecordBatchStreamInput<F = std::future::Ready<TonicResult<FlightRecordBatchSource>>>
57{
58    Ready(FlightRecordBatchSource),
59    Initializer(F),
60}
61
62impl FlightRecordBatchStreamInput {
63    /// Creates an input from a source that is already available.
64    pub fn ready(source: FlightRecordBatchSource) -> Self {
65        Self::Ready(source)
66    }
67}
68
69impl<F> FlightRecordBatchStreamInput<F> {
70    /// Creates an input that obtains its source asynchronously.
71    ///
72    /// Errors from the initializer are returned through the Flight response stream.
73    pub fn initializer(initializer: F) -> Self {
74        Self::Initializer(initializer)
75    }
76}
77
78/// Metrics collector for Flight stream with RAII logging pattern
79struct 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
133/// Coalesces consecutive ready record batches into one outgoing group.
134///
135/// The budgets are soft limits (flush thresholds) rather than strict admission
136/// caps for under-budget batches: `push` appends a batch before reporting whether
137/// to flush, so a group may exceed a budget by up to one batch. A batch that is
138/// itself at or over a budget is never appended: it is forwarded as a singleton,
139/// flushing the current group first when it was encountered inside a group.
140struct BatchAccumulator {
141    batches: Vec<RecordBatch>,
142    rows: usize,
143    bytes: usize,
144}
145
146impl BatchAccumulator {
147    // Sized from the observed upstream batch shape: mito2 commonly emits ~2000-row
148    // batches (~32-94KiB depending on row width), so a 1024-row budget marked every
149    // such batch oversized and coalesced nothing. 4096 lets a few of those batches
150    // group together; the byte budget stays the binding constraint for wide rows.
151    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    /// Whether a batch already reaches a budget on its own, in which case the
164    /// stream forwards it as a singleton instead of coalescing it. This applies
165    /// both to a batch that starts a new group and to a batch encountered inside
166    /// a group (the accumulated group is flushed first).
167    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    /// Appends `batch` and reports whether the group should be flushed now.
176    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        // Append-then-check: the group may exceed a budget by this batch.
181        self.batches.len() >= Self::MAX_BATCHES || Self::reaches_budget(self.rows, self.bytes)
182    }
183
184    /// Removes and returns the accumulated batches in arrival order.
185    ///
186    /// Merging is left to the caller, which owns the stream schema.
187    fn drain(&mut self) -> Vec<RecordBatch> {
188        self.rows = 0;
189        self.bytes = 0;
190        std::mem::take(&mut self.batches)
191    }
192}
193
194/// Forwards ready record batches on the non-verbose path, coalescing
195/// consecutive small batches into one outgoing group before sending.
196struct 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    /// Runs the coalescing loop until the source stream ends or fails. Returns
212    /// `true` on normal EOF, `false` when it stopped early on an error or a failed
213    /// send.
214    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            // The first batch is forwarded immediately, and any batch that is
247            // already at or over a budget on its own passes through as a
248            // singleton. An under-budget batch starts a group instead; batches
249            // encountered inside a group are appended first and the budgets are
250            // checked afterwards.
251            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            // Coalesce ready batches until a budget is reached: under-budget
268            // batches are appended first and the budgets are checked afterwards.
269            // A batch that is itself at or over a budget is held back instead so
270            // it is not copied into the aggregate.
271            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                        // Every fetched batch is counted exactly once here,
283                        // including a batch that is held back below.
284                        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                            // The batch is at or over a budget on its own: flush
291                            // the accumulated group first, then forward it as its
292                            // own singleton (see below) instead of copying it
293                            // into the aggregate.
294                            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                // Reuse the buffer in place: the source batches are dropped here
316                // (and their memory released) instead of staying alive across the
317                // send loop below, which may block on backpressure.
318                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            // A batch encountered inside the group that is itself at or over a
335            // budget was held back: the accumulated group has been flushed and
336            // sent above, so forward it now as its own singleton. Its
337            // total_rows/total_bytes/record_batch_count were already counted when
338            // it was fetched.
339            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                        // Verbose responses preserve their existing per-batch metrics behavior.
447                        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        // Each path reports whether it reached normal EOF. On any error or failed
511        // send the path stops early and the final-metrics tail must be skipped:
512        // final metrics are only sent after a cleanly completed stream.
513        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        // Make the last package pass metrics exactly once at EOF.
532        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    /// On the verbose path, sends every record batch individually and forwards
541    /// partial metrics whenever they change.
542    /// Returns `true` when the source stream reached normal EOF, `false` when it
543    /// stopped early on an error or a failed send.
544    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                            // Safety: `iter` on a type of `Vec1`, which is guaranteed to have
657                            // at least one element.
658                            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        // With capacity 1 the schema is queued and the first-batch send().await
941        // stays pending (the futures mpsc channel reserves a slot per sender, so
942        // the batch may be enqueued but its send flush does not complete) until
943        // the test consumes the schema. This holds the producer at the first-batch
944        // send, where the upstream poll count is observable (an implementation
945        // that polled the source again before sending would have counted 2).
946        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        // No task switch between here and the assertion: the producer is blocked
961        // on the first-batch send, which frees up only once the schema below is
962        // consumed.
963        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    /// Regression test: after a successful merge the original source batches are
1012    /// dropped before the merged group is sent, so their memory is not held across
1013    /// a backpressured send.
1014    ///
1015    /// A `Weak` is used because retention across an await is not observable
1016    /// structurally. The producer is pinned inside the merged send().await: the
1017    /// test only consumes the schema, so the first batch is queued and the merged
1018    /// send flush stays pending (the futures mpsc channel reserves a slot per
1019    /// sender), and the producer can neither complete the merged send nor poll
1020    /// the source again.
1021    #[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        // A primitive (non-buffer-sharing) source array that the test tracks.
1029        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 the only other strong reference: after this, the tracked array is
1040        // owned exclusively by the tracked source batch.
1041        drop(tracked);
1042
1043        // The first batch is forwarded immediately; the tracked batch plus 15 more
1044        // small batches fill a 16-batch group that is merged into one send.
1045        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        // 1 poll for the first batch plus 16 polls for the group.
1067        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        // The merge happens synchronously right after the last fetch, so the source
1075        // batches must be released without waiting for the blocked send.
1076        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        // The at-cap batch is encountered inside a group, but it reaches the row
1165        // budget on its own, so the tiny accumulated group is flushed first and
1166        // the at-cap batch is sent as its own singleton instead of being copied
1167        // into the aggregate. The first batch is always forwarded immediately.
1168        assert_eq!(batches, vec![1, 1, 4096]);
1169        // Three fetches, plus the poll that reports end of stream.
1170        assert_eq!(poll_count.load(std::sync::atomic::Ordering::Relaxed), 4);
1171    }
1172
1173    /// A tiny accumulated group followed by an over-budget batch produces two
1174    /// separate sends (the flushed group, then the over-budget singleton) instead
1175    /// of one aggregate that copies the over-budget batch into the group.
1176    #[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        // [1] is forwarded immediately, [2] + [3] merge into one 2-row group, and
1202        // the at-cap batch is forwarded separately: the aggregate never contains
1203        // the over-budget batch, and it is not copied.
1204        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        // 2100 + 2100 exceeds the 4096-row budget, but the second batch is appended
1233        // before the budget is checked, so the group flushes as one 4200-row batch.
1234        assert_eq!(batches, vec![1, 4200]);
1235        // Three fetches, plus the poll that reports end of stream.
1236        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        // The at-cap batch is current when it is fetched, so it passes through as a
1265        // singleton instead of joining an aggregate.
1266        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    /// A stream that yields one batch then an error, and counts how many times
1302    /// `metrics()` is called. Used to prove the verbose error path does NOT invoke
1303    /// the shared EOF final-metrics tail (which would call `metrics()` one extra
1304    /// time). The public message stream hides the tail after an error, so this
1305    /// producer-side counter is what actually detects the regression.
1306    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    /// On an upstream error the verbose path stops early and skips the shared
1344    /// EOF final-metrics tail: `metrics()` must not be called again after the
1345    /// error. This drives `flight_data_stream` directly and awaits its return,
1346    /// so the producer has fully finished before the counter is read (the public
1347    /// message stream alone would hide the tail and could race with it).
1348    #[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        // should_send_partial_metrics = true selects the verbose path;
1363        // can_send_metrics_before_batch = false disables the heartbeat arm.
1364        FlightRecordBatchStream::flight_data_stream(recordbatches, tx, true, false).await;
1365        // Producer fully returned. Exactly one metrics() call (the per-batch one).
1366        // If the error path fell through to the EOF final-metrics tail, this
1367        // would be 2.
1368        assert_eq!(metrics_calls.load(std::sync::atomic::Ordering::Relaxed), 1);
1369        // The error is the last message; no trailing final-metrics package.
1370        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        // 17 trailing empty batches: the 16-batch budget flushes the first 16 as
1422        // one group, leaving one empty batch for a second group. This exercises the
1423        // batch-count bound rather than relying on EOF to flush.
1424        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        // The at-cap batch is encountered inside a group, so the 1-row group is
1442        // flushed first and the at-cap batch follows as its own singleton. The 17
1443        // trailing empty batches split into two empty groups at the 16-batch budget.
1444        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        // The byte-oversized batch is encountered inside a group, so the leading
1480        // empty batch is flushed first and the oversized batch follows as its own
1481        // singleton instead of being appended to the aggregate.
1482        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        // Decode the merged dictionary to its logical string values by resolving
1524        // each key against the merged dictionary values. Asserting only the
1525        // merged dictionary's values would not catch a key remapping regression
1526        // (keys [0, 0] would still hold equal values but would decode to
1527        // ["first", "first"]). The internal dictionary ordering is an arrow
1528        // implementation detail and is deliberately not asserted.
1529        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}