Skip to main content

common_recordbatch/
lib.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
15#![feature(never_type)]
16
17pub mod adapter;
18pub mod cursor;
19pub mod error;
20pub mod ext;
21pub mod filter;
22pub mod recordbatch;
23pub mod util;
24
25use std::fmt::{self, Write};
26use std::future::Future;
27use std::pin::Pin;
28use std::sync::Arc;
29
30use adapter::RecordBatchMetrics;
31use arc_swap::ArcSwapOption;
32use common_base::readable_size::ReadableSize;
33use common_error::ext::BoxedError;
34use common_memory_manager::{
35    MemoryGuard, MemoryManager, MemoryMetrics, OnExhaustedPolicy, PermitGranularity,
36};
37use common_telemetry::tracing::Span;
38pub use datafusion::physical_plan::SendableRecordBatchStream as DfSendableRecordBatchStream;
39use datatypes::arrow::array::{Array, ArrayRef, AsArray, StringBuilder};
40use datatypes::arrow::compute::SortOptions;
41use datatypes::arrow::datatypes::{DataType as ArrowDataType, Field};
42use datatypes::arrow::error::ArrowError;
43pub use datatypes::arrow::record_batch::RecordBatch as DfRecordBatch;
44use datatypes::arrow::util::display::{
45    ArrayFormatter, ArrayFormatterFactory, DisplayIndex, FormatOptions, FormatResult,
46};
47use datatypes::arrow::util::pretty::pretty_format_batches_with_options;
48use datatypes::extension::json::is_json_extension_type;
49use datatypes::prelude::{ConcreteDataType, DataType, VectorRef};
50use datatypes::schema::{ColumnSchema, Schema, SchemaRef};
51use datatypes::types::{JsonFormat, StructField, StructType, jsonb_to_string};
52use error::Result;
53use futures::task::{Context, Poll};
54use futures::{Stream, TryStreamExt};
55pub use recordbatch::RecordBatch;
56use snafu::{IntoError, ResultExt, ensure};
57
58use crate::error::{ArrowComputeSnafu, NewDfRecordBatchSnafu};
59
60pub trait RecordBatchStream: Stream<Item = Result<RecordBatch>> {
61    fn name(&self) -> &str {
62        "RecordBatchStream"
63    }
64
65    fn schema(&self) -> SchemaRef;
66
67    fn output_ordering(&self) -> Option<&[OrderOption]>;
68
69    fn metrics(&self) -> Option<RecordBatchMetrics>;
70}
71
72pub type SendableRecordBatchStream = Pin<Box<dyn RecordBatchStream + Send>>;
73
74#[derive(Debug, Clone, PartialEq, Eq)]
75pub struct OrderOption {
76    pub name: String,
77    pub options: SortOptions,
78}
79
80/// A wrapper that maps a [RecordBatchStream] to a new [RecordBatchStream] by applying a function to each [RecordBatch].
81///
82/// The mapper function is applied to each [RecordBatch] in the stream.
83/// The schema of the new [RecordBatchStream] is the same as the schema of the inner [RecordBatchStream] after applying the schema mapper function.
84/// The output ordering of the new [RecordBatchStream] is the same as the output ordering of the inner [RecordBatchStream].
85/// The metrics of the new [RecordBatchStream] is the same as the metrics of the inner [RecordBatchStream] if it is not `None`.
86pub struct SendableRecordBatchMapper {
87    inner: SendableRecordBatchStream,
88    /// The mapper function is applied to each [RecordBatch] in the stream.
89    /// The original schema and the mapped schema are passed to the mapper function.
90    mapper: fn(RecordBatch, &SchemaRef, &SchemaRef) -> Result<RecordBatch>,
91    /// The schema of the new [RecordBatchStream] is the same as the schema of the inner [RecordBatchStream] after applying the schema mapper function.
92    schema: SchemaRef,
93    /// Whether the mapper function is applied to each [RecordBatch] in the stream.
94    apply_mapper: bool,
95}
96
97/// Maps the json type to string in the batch.
98///
99/// The json type is mapped to string by converting the json value to string.
100/// The batch is updated to have the same number of columns as the original batch,
101/// but with the json type mapped to string.
102pub fn map_json_type_to_string(
103    batch: RecordBatch,
104    original_schema: &SchemaRef,
105    mapped_schema: &SchemaRef,
106) -> Result<RecordBatch> {
107    let mut vectors = Vec::with_capacity(original_schema.column_schemas().len());
108    for (vector, schema) in batch.columns().iter().zip(original_schema.column_schemas()) {
109        if let ConcreteDataType::Json(j) = &schema.data_type {
110            if matches!(&j.format, JsonFormat::Jsonb) {
111                let mut string_vector_builder = StringBuilder::new();
112                let binary_vector = vector.as_binary::<i32>();
113                for value in binary_vector.iter() {
114                    let Some(value) = value else {
115                        string_vector_builder.append_null();
116                        continue;
117                    };
118                    let string_value =
119                        jsonb_to_string(value).with_context(|_| error::CastVectorSnafu {
120                            from_type: schema.data_type.clone(),
121                            to_type: ConcreteDataType::string_datatype(),
122                        })?;
123                    string_vector_builder.append_value(string_value);
124                }
125
126                let string_vector = string_vector_builder.finish();
127                vectors.push(Arc::new(string_vector) as ArrayRef);
128            } else {
129                vectors.push(vector.clone());
130            }
131        } else {
132            vectors.push(vector.clone());
133        }
134    }
135
136    let record_batch = datatypes::arrow::record_batch::RecordBatch::try_new(
137        mapped_schema.arrow_schema().clone(),
138        vectors,
139    )
140    .context(NewDfRecordBatchSnafu)?;
141    Ok(RecordBatch::from_df_record_batch(
142        mapped_schema.clone(),
143        record_batch,
144    ))
145}
146
147/// Maps the json type to string in the schema.
148///
149/// The json type is mapped to string by converting the json value to string.
150/// The schema is updated to have the same number of columns as the original schema,
151/// but with the json type mapped to string.
152///
153/// Returns the new schema and whether the schema needs to be mapped to string.
154pub fn map_json_type_to_string_schema(schema: SchemaRef) -> (SchemaRef, bool) {
155    let mut new_columns = Vec::with_capacity(schema.column_schemas().len());
156    let mut apply_mapper = false;
157    for column in schema.column_schemas() {
158        if matches!(column.data_type, ConcreteDataType::Json(_)) {
159            new_columns.push(ColumnSchema::new(
160                column.name.clone(),
161                ConcreteDataType::string_datatype(),
162                column.is_nullable(),
163            ));
164            apply_mapper = true;
165        } else {
166            new_columns.push(column.clone());
167        }
168    }
169    (Arc::new(Schema::new(new_columns)), apply_mapper)
170}
171
172/// Replaces dictionary types with their value types, including dictionaries nested in lists and
173/// structs.
174pub fn map_dictionary_to_values_data_type(data_type: &ConcreteDataType) -> ConcreteDataType {
175    match data_type {
176        ConcreteDataType::Dictionary(dictionary) => {
177            map_dictionary_to_values_data_type(dictionary.value_type())
178        }
179        ConcreteDataType::List(list) => ConcreteDataType::list_datatype(Arc::new(
180            map_dictionary_to_values_data_type(list.item_type()),
181        )),
182        ConcreteDataType::Struct(struct_type) => {
183            let fields = struct_type
184                .fields()
185                .iter()
186                .map(|field| {
187                    StructField::new(
188                        field.name(),
189                        map_dictionary_to_values_data_type(field.data_type()),
190                        field.is_nullable(),
191                    )
192                })
193                .collect();
194            ConcreteDataType::struct_datatype(StructType::new(Arc::new(fields)))
195        }
196        _ => data_type.clone(),
197    }
198}
199
200/// Maps dictionary columns in a schema to their value types.
201pub fn map_dictionary_to_values_schema(schema: SchemaRef) -> (SchemaRef, bool) {
202    let mut apply_mapper = false;
203    let columns: Vec<_> = schema
204        .column_schemas()
205        .iter()
206        .map(|column| {
207            let data_type = map_dictionary_to_values_data_type(&column.data_type);
208            apply_mapper |= data_type != column.data_type;
209            let mut column = column.clone();
210            column.data_type = data_type;
211            column
212        })
213        .collect();
214
215    if !apply_mapper {
216        return (schema, false);
217    }
218
219    // Query projections may contain duplicate column names, which SchemaBuilder rejects. Preserve
220    // the existing Arrow fields and only replace their data types so field and schema metadata are
221    // retained as well.
222    let fields = schema
223        .arrow_schema()
224        .fields()
225        .iter()
226        .zip(&columns)
227        .map(|(field, column)| {
228            Arc::new(
229                field
230                    .as_ref()
231                    .clone()
232                    .with_data_type(column.data_type.as_arrow_type()),
233            )
234        })
235        .collect::<Vec<_>>();
236    let arrow_schema =
237        datatypes::arrow::datatypes::Schema::new(fields).with_metadata(schema.metadata().clone());
238    (
239        Arc::new(Schema::try_from(Arc::new(arrow_schema)).unwrap()),
240        true,
241    )
242}
243
244/// Expands dictionary arrays to their value arrays according to `mapped_schema`.
245pub fn map_dictionary_to_values(
246    batch: RecordBatch,
247    original_schema: &SchemaRef,
248    mapped_schema: &SchemaRef,
249) -> Result<RecordBatch> {
250    let arrays = batch
251        .columns()
252        .iter()
253        .zip(original_schema.column_schemas())
254        .zip(mapped_schema.column_schemas())
255        .map(|((array, original), mapped)| {
256            if original.data_type == mapped.data_type {
257                Ok(array.clone())
258            } else {
259                datatypes::arrow::compute::cast(array, &mapped.data_type.as_arrow_type())
260                    .context(ArrowComputeSnafu)
261            }
262        })
263        .collect::<Result<Vec<_>>>()?;
264
265    let record_batch = DfRecordBatch::try_new(mapped_schema.arrow_schema().clone(), arrays)
266        .context(NewDfRecordBatchSnafu)?;
267    Ok(RecordBatch::from_df_record_batch(
268        mapped_schema.clone(),
269        record_batch,
270    ))
271}
272
273impl SendableRecordBatchMapper {
274    /// Creates a new [SendableRecordBatchMapper] with the given inner [RecordBatchStream], mapper function, and schema mapper function.
275    pub fn new(
276        inner: SendableRecordBatchStream,
277        mapper: fn(RecordBatch, &SchemaRef, &SchemaRef) -> Result<RecordBatch>,
278        schema_mapper: fn(SchemaRef) -> (SchemaRef, bool),
279    ) -> Self {
280        let (mapped_schema, apply_mapper) = schema_mapper(inner.schema());
281        Self {
282            inner,
283            mapper,
284            schema: mapped_schema,
285            apply_mapper,
286        }
287    }
288}
289
290impl RecordBatchStream for SendableRecordBatchMapper {
291    fn name(&self) -> &str {
292        "SendableRecordBatchMapper"
293    }
294
295    fn schema(&self) -> SchemaRef {
296        self.schema.clone()
297    }
298
299    fn output_ordering(&self) -> Option<&[OrderOption]> {
300        self.inner.output_ordering()
301    }
302
303    fn metrics(&self) -> Option<RecordBatchMetrics> {
304        self.inner.metrics()
305    }
306}
307
308impl Stream for SendableRecordBatchMapper {
309    type Item = Result<RecordBatch>;
310
311    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
312        if self.apply_mapper {
313            Pin::new(&mut self.inner).poll_next(cx).map(|opt| {
314                opt.map(|result| {
315                    result
316                        .and_then(|batch| (self.mapper)(batch, &self.inner.schema(), &self.schema))
317                })
318            })
319        } else {
320            Pin::new(&mut self.inner).poll_next(cx)
321        }
322    }
323}
324
325/// EmptyRecordBatchStream can be used to create a RecordBatchStream
326/// that will produce no results
327pub struct EmptyRecordBatchStream {
328    /// Schema wrapped by Arc
329    schema: SchemaRef,
330}
331
332impl EmptyRecordBatchStream {
333    /// Create an empty RecordBatchStream
334    pub fn new(schema: SchemaRef) -> Self {
335        Self { schema }
336    }
337}
338
339impl RecordBatchStream for EmptyRecordBatchStream {
340    fn schema(&self) -> SchemaRef {
341        self.schema.clone()
342    }
343
344    fn output_ordering(&self) -> Option<&[OrderOption]> {
345        None
346    }
347
348    fn metrics(&self) -> Option<RecordBatchMetrics> {
349        None
350    }
351}
352
353impl Stream for EmptyRecordBatchStream {
354    type Item = Result<RecordBatch>;
355
356    fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
357        Poll::Ready(None)
358    }
359}
360
361#[derive(Debug, PartialEq)]
362pub struct RecordBatches {
363    schema: SchemaRef,
364    batches: Vec<RecordBatch>,
365}
366
367impl RecordBatches {
368    pub fn try_from_columns<I: IntoIterator<Item = VectorRef>>(
369        schema: SchemaRef,
370        columns: I,
371    ) -> Result<Self> {
372        let batches = vec![RecordBatch::new(schema.clone(), columns)?];
373        Ok(Self { schema, batches })
374    }
375
376    pub async fn try_collect(stream: SendableRecordBatchStream) -> Result<Self> {
377        let schema = stream.schema();
378        let batches = stream.try_collect::<Vec<_>>().await?;
379        Ok(Self { schema, batches })
380    }
381
382    #[inline]
383    pub fn empty() -> Self {
384        Self {
385            schema: Arc::new(Schema::new(vec![])),
386            batches: vec![],
387        }
388    }
389
390    pub fn iter(&self) -> impl Iterator<Item = &RecordBatch> {
391        self.batches.iter()
392    }
393
394    pub fn pretty_print(&self) -> Result<String> {
395        let df_batches = &self
396            .iter()
397            .map(|x| x.df_record_batch().clone())
398            .collect::<Vec<_>>();
399        let options =
400            FormatOptions::default().with_formatter_factory(Some(&BinaryFormatterFactory));
401        let result =
402            pretty_format_batches_with_options(df_batches, &options).context(error::FormatSnafu)?;
403
404        Ok(result.to_string())
405    }
406
407    pub fn try_new(schema: SchemaRef, batches: Vec<RecordBatch>) -> Result<Self> {
408        for batch in &batches {
409            ensure!(
410                batch.schema == schema,
411                error::CreateRecordBatchesSnafu {
412                    reason: format!(
413                        "expect RecordBatch schema equals {:?}, actual: {:?}",
414                        schema, batch.schema
415                    )
416                }
417            )
418        }
419        Ok(Self { schema, batches })
420    }
421
422    pub fn schema(&self) -> SchemaRef {
423        self.schema.clone()
424    }
425
426    pub fn take(self) -> Vec<RecordBatch> {
427        self.batches
428    }
429
430    pub fn as_stream(&self) -> SendableRecordBatchStream {
431        Box::pin(SimpleRecordBatchStream {
432            inner: RecordBatches {
433                schema: self.schema(),
434                batches: self.batches.clone(),
435            },
436            index: 0,
437        })
438    }
439}
440
441#[derive(Debug)]
442struct BinaryFormatterFactory;
443
444impl ArrayFormatterFactory for BinaryFormatterFactory {
445    fn create_array_formatter<'a>(
446        &self,
447        array: &'a dyn Array,
448        options: &FormatOptions<'a>,
449        field: Option<&'a Field>,
450    ) -> std::result::Result<Option<ArrayFormatter<'a>>, ArrowError> {
451        if !array.data_type().is_binary() {
452            return Ok(None);
453        }
454
455        Ok(Some(ArrayFormatter::new(
456            Box::new(BinaryFormatter {
457                array,
458                is_json: field.is_some_and(is_json_extension_type),
459                default: ArrayFormatter::try_new(array, options)?,
460                null: options.null(),
461            }),
462            options.safe(),
463        )))
464    }
465}
466
467struct BinaryFormatter<'a> {
468    array: &'a dyn Array,
469    is_json: bool,
470    default: ArrayFormatter<'a>,
471    null: &'a str,
472}
473
474impl DisplayIndex for BinaryFormatter<'_> {
475    fn write(&self, idx: usize, f: &mut dyn Write) -> FormatResult {
476        if !self.is_json {
477            self.default.value(idx).write(f)?;
478            return Ok(());
479        }
480
481        if self.array.is_null(idx) {
482            write!(f, "{}", self.null)?;
483        } else {
484            let bytes = match self.array.data_type() {
485                ArrowDataType::Binary => self.array.as_binary::<i32>().value(idx),
486                ArrowDataType::LargeBinary => self.array.as_binary::<i64>().value(idx),
487                ArrowDataType::BinaryView => self.array.as_binary_view().value(idx),
488                _ => unreachable!(),
489            };
490            let value =
491                jsonb_to_string(bytes).map_err(|e| ArrowError::ExternalError(Box::new(e)))?;
492            write!(f, "{value}")?;
493        }
494        Ok(())
495    }
496}
497
498impl IntoIterator for RecordBatches {
499    type Item = RecordBatch;
500    type IntoIter = std::vec::IntoIter<Self::Item>;
501
502    fn into_iter(self) -> Self::IntoIter {
503        self.batches.into_iter()
504    }
505}
506
507pub struct SimpleRecordBatchStream {
508    inner: RecordBatches,
509    index: usize,
510}
511
512impl RecordBatchStream for SimpleRecordBatchStream {
513    fn schema(&self) -> SchemaRef {
514        self.inner.schema()
515    }
516
517    fn output_ordering(&self) -> Option<&[OrderOption]> {
518        None
519    }
520
521    fn metrics(&self) -> Option<RecordBatchMetrics> {
522        None
523    }
524}
525
526impl Stream for SimpleRecordBatchStream {
527    type Item = Result<RecordBatch>;
528
529    fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
530        Poll::Ready(if self.index < self.inner.batches.len() {
531            let batch = self.inner.batches[self.index].clone();
532            self.index += 1;
533            Some(Ok(batch))
534        } else {
535            None
536        })
537    }
538}
539
540/// Adapt a [Stream] of [RecordBatch] to a [RecordBatchStream].
541pub struct RecordBatchStreamWrapper<S> {
542    pub schema: SchemaRef,
543    pub stream: S,
544    pub output_ordering: Option<Vec<OrderOption>>,
545    pub metrics: Arc<ArcSwapOption<RecordBatchMetrics>>,
546    pub span: Span,
547}
548
549impl<S> RecordBatchStreamWrapper<S> {
550    /// Creates a [RecordBatchStreamWrapper] without output ordering requirement.
551    pub fn new(schema: SchemaRef, stream: S) -> RecordBatchStreamWrapper<S> {
552        RecordBatchStreamWrapper {
553            schema,
554            stream,
555            output_ordering: None,
556            metrics: Default::default(),
557            span: Span::current(),
558        }
559    }
560}
561
562impl<S: Stream<Item = Result<RecordBatch>> + Unpin> RecordBatchStream
563    for RecordBatchStreamWrapper<S>
564{
565    fn name(&self) -> &str {
566        "RecordBatchStreamWrapper"
567    }
568
569    fn schema(&self) -> SchemaRef {
570        self.schema.clone()
571    }
572
573    fn output_ordering(&self) -> Option<&[OrderOption]> {
574        self.output_ordering.as_deref()
575    }
576
577    fn metrics(&self) -> Option<RecordBatchMetrics> {
578        self.metrics.load().as_ref().map(|s| s.as_ref().clone())
579    }
580}
581
582impl<S: Stream<Item = Result<RecordBatch>> + Unpin> Stream for RecordBatchStreamWrapper<S> {
583    type Item = Result<RecordBatch>;
584
585    fn poll_next(mut self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
586        let _entered = self.span.clone().entered();
587        Pin::new(&mut self.stream).poll_next(ctx)
588    }
589}
590
591/// Memory tracker for RecordBatch streams. Clone to share the same limit across queries.
592///
593/// Each stream acquires quota independently from this tracker.
594#[derive(Clone)]
595pub struct QueryMemoryTracker {
596    manager: MemoryManager<CallbackMemoryMetrics>,
597    metrics: CallbackMemoryMetrics,
598    on_exhausted_policy: OnExhaustedPolicy,
599}
600
601impl fmt::Debug for QueryMemoryTracker {
602    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
603        f.debug_struct("QueryMemoryTracker")
604            .field("current", &self.current())
605            .field("limit", &self.limit())
606            .field("on_exhausted_policy", &self.on_exhausted_policy)
607            .field("on_update", &self.metrics.has_on_update())
608            .field("on_exhausted", &self.metrics.has_on_exhausted())
609            .field("on_rejected", &self.metrics.has_on_rejected())
610            .finish()
611    }
612}
613
614impl QueryMemoryTracker {
615    /// Create a builder for a query memory tracker.
616    pub fn builder(
617        limit: usize,
618        on_exhausted_policy: OnExhaustedPolicy,
619    ) -> QueryMemoryTrackerBuilder {
620        QueryMemoryTrackerBuilder {
621            limit,
622            on_exhausted_policy,
623            on_update: None,
624            on_exhausted: None,
625            on_reject: None,
626        }
627    }
628
629    fn new_stream_tracker(&self) -> StreamMemoryTracker {
630        StreamMemoryTracker {
631            tracker: self.clone(),
632            guard: self.manager.try_acquire(0).unwrap(),
633            tracked_bytes: 0,
634        }
635    }
636    /// Get the current memory usage in bytes.
637    pub fn current(&self) -> usize {
638        self.manager.used_bytes() as usize
639    }
640
641    fn limit(&self) -> usize {
642        self.manager.limit_bytes() as usize
643    }
644
645    fn reject_error(
646        &self,
647        current: usize,
648        additional: usize,
649        stream_tracked: usize,
650    ) -> error::Error {
651        let limit = self.limit();
652        let msg = format!(
653            "{} requested, {} used globally ({}%), {} used by this stream, hard limit: {}",
654            ReadableSize(additional as u64),
655            ReadableSize(current as u64),
656            (current * 100).checked_div(limit).unwrap_or(0),
657            ReadableSize(stream_tracked as u64),
658            ReadableSize(limit as u64)
659        );
660        error::ExceedMemoryLimitSnafu { msg }.build()
661    }
662
663    fn inc_rejected(&self) {
664        self.metrics.inc_rejected();
665    }
666}
667
668/// Builder for constructing a [`QueryMemoryTracker`] with optional callbacks.
669pub struct QueryMemoryTrackerBuilder {
670    limit: usize,
671    on_exhausted_policy: OnExhaustedPolicy,
672    on_update: Option<UpdateCallback>,
673    on_exhausted: Option<UnitCallback>,
674    on_reject: Option<RejectCallback>,
675}
676
677impl QueryMemoryTrackerBuilder {
678    /// Set a callback to be called whenever the usage changes successfully.
679    /// The callback receives the new total usage in bytes.
680    ///
681    /// # Note
682    /// The callback is called after both successful `track()` and stream drop.
683    /// Usage is exact in unlimited mode and 1KB-aligned in limited mode.
684    pub fn on_update<F>(mut self, on_update: F) -> Self
685    where
686        F: Fn(usize) + Send + Sync + 'static,
687    {
688        self.on_update = Some(Arc::new(on_update));
689        self
690    }
691
692    /// Set a callback to be called when memory is unavailable for immediate acquisition.
693    ///
694    /// # Note
695    /// This is called when the non-blocking allocation fast path fails.
696    /// Requests using `OnExhaustedPolicy::Wait` may still succeed after waiting.
697    /// It is never called when `limit == 0` (unlimited mode).
698    pub fn on_exhausted<F>(mut self, on_exhausted: F) -> Self
699    where
700        F: Fn() + Send + Sync + 'static,
701    {
702        self.on_exhausted = Some(Arc::new(on_exhausted));
703        self
704    }
705
706    /// Set a callback to be called when the request ultimately fails due to memory pressure.
707    pub fn on_reject<F>(mut self, on_reject: F) -> Self
708    where
709        F: Fn() + Send + Sync + 'static,
710    {
711        self.on_reject = Some(Arc::new(on_reject));
712        self
713    }
714
715    /// Build a [`QueryMemoryTracker`] from this builder.
716    pub fn build(self) -> QueryMemoryTracker {
717        let metrics = CallbackMemoryMetrics::new(self.on_update, self.on_exhausted, self.on_reject);
718        let manager = MemoryManager::with_granularity(
719            self.limit as u64,
720            PermitGranularity::Kilobyte,
721            metrics.clone(),
722        );
723
724        QueryMemoryTracker {
725            manager,
726            metrics,
727            on_exhausted_policy: self.on_exhausted_policy,
728        }
729    }
730}
731
732struct StreamMemoryTracker {
733    tracker: QueryMemoryTracker,
734    guard: MemoryGuard<CallbackMemoryMetrics>,
735    tracked_bytes: usize,
736}
737
738type MemoryAcquireResult = std::result::Result<(), common_memory_manager::Error>;
739
740impl StreamMemoryTracker {
741    fn inc_rejected(&self) {
742        self.tracker.inc_rejected();
743    }
744
745    fn try_track(&mut self, additional: usize) -> Result<()> {
746        if self.guard.try_acquire_additional(additional as u64) {
747            self.tracked_bytes = self.tracked_bytes.saturating_add(additional);
748            Ok(())
749        } else {
750            Err(self.reject_error(additional))
751        }
752    }
753
754    async fn track_with_policy(mut self, additional: usize) -> (Self, MemoryAcquireResult) {
755        let result = self
756            .guard
757            .acquire_additional_with_policy(additional as u64, self.tracker.on_exhausted_policy)
758            .await;
759        if result.is_ok() {
760            self.tracked_bytes = self.tracked_bytes.saturating_add(additional);
761        }
762        (self, result)
763    }
764
765    fn reject_error(&self, additional: usize) -> error::Error {
766        let current = self.tracker.current();
767        self.tracker
768            .reject_error(current, additional, self.tracked_bytes)
769    }
770
771    fn wait_error(&self, additional: usize, source: common_memory_manager::Error) -> error::Error {
772        match source {
773            common_memory_manager::Error::MemoryLimitExceeded { .. } => {
774                self.reject_error(additional)
775            }
776            common_memory_manager::Error::MemoryAcquireTimeout { waited, .. } => {
777                let current = self.tracker.current();
778                let limit = self.tracker.limit();
779                let msg = format!(
780                    "timed out waiting {:?} for {}, {} used globally ({}%), {} used by this stream, hard limit: {}",
781                    waited,
782                    ReadableSize(additional as u64),
783                    ReadableSize(current as u64),
784                    (current * 100).checked_div(limit).unwrap_or(0),
785                    ReadableSize(self.tracked_bytes as u64),
786                    ReadableSize(limit as u64)
787                );
788                error::ExceedMemoryLimitSnafu { msg }.build()
789            }
790            error => error::ExternalSnafu.into_error(BoxedError::new(error)),
791        }
792    }
793}
794
795type PendingTrackFuture = Pin<
796    Box<dyn Future<Output = (StreamMemoryTracker, RecordBatch, usize, MemoryAcquireResult)> + Send>,
797>;
798
799#[derive(Clone)]
800struct CallbackMemoryMetrics {
801    inner: Arc<CallbackMemoryMetricsInner>,
802}
803
804type UpdateCallback = Arc<dyn Fn(usize) + Send + Sync>;
805type UnitCallback = Arc<dyn Fn() + Send + Sync>;
806type RejectCallback = UnitCallback;
807
808struct CallbackMemoryMetricsInner {
809    on_update: Option<UpdateCallback>,
810    on_exhausted: Option<UnitCallback>,
811    on_reject: Option<RejectCallback>,
812}
813
814impl CallbackMemoryMetrics {
815    fn new(
816        on_update: Option<UpdateCallback>,
817        on_exhausted: Option<UnitCallback>,
818        on_reject: Option<RejectCallback>,
819    ) -> Self {
820        Self {
821            inner: Arc::new(CallbackMemoryMetricsInner {
822                on_update,
823                on_exhausted,
824                on_reject,
825            }),
826        }
827    }
828
829    fn has_on_update(&self) -> bool {
830        self.inner.on_update.is_some()
831    }
832
833    fn has_on_exhausted(&self) -> bool {
834        self.inner.on_exhausted.is_some()
835    }
836
837    fn has_on_rejected(&self) -> bool {
838        self.inner.on_reject.is_some()
839    }
840
841    fn inc_rejected(&self) {
842        if let Some(callback) = &self.inner.on_reject {
843            callback();
844        }
845    }
846}
847
848impl MemoryMetrics for CallbackMemoryMetrics {
849    fn set_limit(&self, _: i64) {}
850
851    fn set_in_use(&self, bytes: i64) {
852        if let Some(callback) = &self.inner.on_update {
853            callback(bytes.max(0) as usize);
854        }
855    }
856
857    fn inc_exhausted(&self, _: &str) {
858        if let Some(callback) = &self.inner.on_exhausted {
859            callback();
860        }
861    }
862}
863
864/// A wrapper stream that tracks memory usage of RecordBatches.
865pub struct MemoryTrackedStream {
866    inner: SendableRecordBatchStream,
867    tracker: Option<StreamMemoryTracker>,
868    // Waiting stores a batch that has already been pulled from the inner stream but has not yet
869    // acquired additional quota. This keeps `poll_next()` non-blocking and allows bounded waits,
870    // at the cost of temporarily holding one untracked batch per blocked stream in memory.
871    waiting: Option<PendingTrackFuture>,
872}
873
874impl MemoryTrackedStream {
875    pub fn new(inner: SendableRecordBatchStream, tracker: QueryMemoryTracker) -> Self {
876        Self {
877            inner,
878            tracker: Some(tracker.new_stream_tracker()),
879            waiting: None,
880        }
881    }
882
883    fn ready_tracker_mut(&mut self) -> &mut StreamMemoryTracker {
884        debug_assert!(
885            self.waiting.is_none(),
886            "a ready tracker must not coexist with a waiting future"
887        );
888        self.tracker.as_mut().unwrap()
889    }
890
891    fn enter_waiting(&mut self, batch: RecordBatch, additional: usize) {
892        debug_assert!(
893            self.waiting.is_none(),
894            "enter_waiting should only be called from the ready state"
895        );
896        debug_assert!(
897            self.tracker.is_some(),
898            "enter_waiting requires a tracker in the ready state"
899        );
900        let tracker = self.tracker.take().unwrap();
901        self.waiting = Some(Self::start_waiting(tracker, batch, additional));
902    }
903
904    fn start_waiting(
905        tracker: StreamMemoryTracker,
906        batch: RecordBatch,
907        additional: usize,
908    ) -> PendingTrackFuture {
909        Box::pin(async move {
910            let (tracker, result) = tracker.track_with_policy(additional).await;
911            (tracker, batch, additional, result)
912        })
913    }
914
915    fn poll_waiting(&mut self, cx: &mut Context<'_>) -> Poll<Option<Result<RecordBatch>>> {
916        let future = self.waiting.as_mut().unwrap();
917        match future.as_mut().poll(cx) {
918            Poll::Ready((tracker, batch, additional, result)) => {
919                let output = match result {
920                    Ok(()) => Ok(batch),
921                    Err(error) => {
922                        tracker.inc_rejected();
923                        Err(tracker.wait_error(additional, error))
924                    }
925                };
926                self.waiting = None;
927                self.tracker = Some(tracker);
928                Poll::Ready(Some(output))
929            }
930            Poll::Pending => Poll::Pending,
931        }
932    }
933
934    fn poll_batch(
935        &mut self,
936        batch: RecordBatch,
937        cx: &mut Context<'_>,
938    ) -> Poll<Option<Result<RecordBatch>>> {
939        let additional = batch.logical_slice_memory_size();
940        let tracker = self.ready_tracker_mut();
941
942        if let Err(error) = tracker.try_track(additional) {
943            match tracker.tracker.on_exhausted_policy {
944                OnExhaustedPolicy::Fail => {
945                    tracker.inc_rejected();
946                    return Poll::Ready(Some(Err(error)));
947                }
948                // `Wait` is a deliberate tradeoff: the batch has already been materialized, so we
949                // keep it in memory while waiting for quota instead of failing immediately. Under
950                // contention, real memory usage can therefore exceed `scan_memory_limit` by up to
951                // one buffered batch per blocked stream.
952                OnExhaustedPolicy::Wait { .. } => {
953                    self.enter_waiting(batch, additional);
954                    return self.poll_waiting(cx);
955                }
956            }
957        }
958
959        Poll::Ready(Some(Ok(batch)))
960    }
961}
962
963impl Stream for MemoryTrackedStream {
964    type Item = Result<RecordBatch>;
965
966    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
967        if self.waiting.is_some() {
968            return self.poll_waiting(cx);
969        }
970
971        match Pin::new(&mut self.inner).poll_next(cx) {
972            Poll::Ready(Some(Ok(batch))) => self.poll_batch(batch, cx),
973            Poll::Ready(Some(Err(error))) => Poll::Ready(Some(Err(error))),
974            Poll::Ready(None) => Poll::Ready(None),
975            Poll::Pending => Poll::Pending,
976        }
977    }
978
979    fn size_hint(&self) -> (usize, Option<usize>) {
980        self.inner.size_hint()
981    }
982}
983
984impl RecordBatchStream for MemoryTrackedStream {
985    fn schema(&self) -> SchemaRef {
986        self.inner.schema()
987    }
988
989    fn output_ordering(&self) -> Option<&[OrderOption]> {
990        self.inner.output_ordering()
991    }
992
993    fn metrics(&self) -> Option<RecordBatchMetrics> {
994        self.inner.metrics()
995    }
996}
997
998#[cfg(test)]
999mod tests {
1000    use std::sync::Arc;
1001    use std::sync::atomic::{AtomicUsize, Ordering};
1002    use std::time::Duration;
1003
1004    use common_memory_manager::{OnExhaustedPolicy, PermitGranularity};
1005    use datatypes::arrow::array::{
1006        DictionaryArray, Int32Array, ListArray, StringArray, UInt32Array,
1007    };
1008    use datatypes::arrow::buffer::OffsetBuffer;
1009    use datatypes::arrow::datatypes::{
1010        DataType as ArrowDataType, Field, Int32Type, Schema as ArrowSchema,
1011    };
1012    use datatypes::prelude::{ConcreteDataType, VectorRef};
1013    use datatypes::schema::{ColumnSchema, Schema};
1014    use datatypes::vectors::{BooleanVector, Int32Vector, StringVector};
1015    use futures::StreamExt;
1016    use tokio::time::{sleep, timeout};
1017
1018    use super::*;
1019
1020    fn large_string_batch(bytes: usize) -> RecordBatch {
1021        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1022            "payload",
1023            ConcreteDataType::string_datatype(),
1024            false,
1025        )]));
1026        let payload = "x".repeat(bytes);
1027        let vector: VectorRef = Arc::new(StringVector::from(vec![payload]));
1028        RecordBatch::new(schema, vec![vector]).unwrap()
1029    }
1030
1031    fn aligned_tracked_bytes(bytes: usize) -> usize {
1032        PermitGranularity::Kilobyte
1033            .permits_to_bytes(PermitGranularity::Kilobyte.bytes_to_permits(bytes as u64))
1034            as usize
1035    }
1036
1037    #[tokio::test]
1038    async fn test_memory_tracked_stream_charges_logical_slice_size() {
1039        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1040            "payload",
1041            ConcreteDataType::string_datatype(),
1042            false,
1043        )]));
1044        let payloads: Vec<_> = (0..1024)
1045            .map(|value| format!("payload-{value:04}"))
1046            .collect();
1047        let batch = RecordBatch::new(schema, vec![Arc::new(StringVector::from(payloads)) as _])
1048            .unwrap()
1049            .slice(512, 1)
1050            .unwrap();
1051        let expected_bytes = aligned_tracked_bytes(batch.logical_slice_memory_size());
1052        assert!(expected_bytes < aligned_tracked_bytes(batch.buffer_memory_size()));
1053        let tracker = QueryMemoryTracker::builder(MB, OnExhaustedPolicy::Fail).build();
1054        let mut stream = MemoryTrackedStream::new(
1055            RecordBatches::try_new(batch.schema.clone(), vec![batch])
1056                .unwrap()
1057                .as_stream(),
1058            tracker.clone(),
1059        );
1060
1061        stream.next().await.unwrap().unwrap();
1062
1063        assert_eq!(tracker.current(), expected_bytes);
1064    }
1065
1066    #[test]
1067    fn test_recordbatches_try_from_columns() {
1068        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1069            "a",
1070            ConcreteDataType::int32_datatype(),
1071            false,
1072        )]));
1073        let result = RecordBatches::try_from_columns(
1074            schema.clone(),
1075            vec![Arc::new(StringVector::from(vec!["hello", "world"])) as _],
1076        );
1077        assert!(result.is_err());
1078
1079        let v: VectorRef = Arc::new(Int32Vector::from_slice([1, 2]));
1080        let expected = vec![RecordBatch::new(schema.clone(), vec![v.clone()]).unwrap()];
1081        let r = RecordBatches::try_from_columns(schema, vec![v]).unwrap();
1082        assert_eq!(r.take(), expected);
1083    }
1084
1085    #[test]
1086    fn test_recordbatches_try_new() {
1087        let column_a = ColumnSchema::new("a", ConcreteDataType::int32_datatype(), false);
1088        let column_b = ColumnSchema::new("b", ConcreteDataType::string_datatype(), false);
1089        let column_c = ColumnSchema::new("c", ConcreteDataType::boolean_datatype(), false);
1090
1091        let va: VectorRef = Arc::new(Int32Vector::from_slice([1, 2]));
1092        let vb: VectorRef = Arc::new(StringVector::from(vec!["hello", "world"]));
1093        let vc: VectorRef = Arc::new(BooleanVector::from(vec![true, false]));
1094
1095        let schema1 = Arc::new(Schema::new(vec![column_a.clone(), column_b]));
1096        let batch1 = RecordBatch::new(schema1.clone(), vec![va.clone(), vb]).unwrap();
1097
1098        let schema2 = Arc::new(Schema::new(vec![column_a, column_c]));
1099        let batch2 = RecordBatch::new(schema2.clone(), vec![va, vc]).unwrap();
1100
1101        let result = RecordBatches::try_new(schema1.clone(), vec![batch1.clone(), batch2]);
1102        assert!(result.is_err());
1103        assert_eq!(
1104            result.unwrap_err().to_string(),
1105            format!(
1106                "Failed to create RecordBatches, reason: expect RecordBatch schema equals {schema1:?}, actual: {schema2:?}",
1107            )
1108        );
1109
1110        let batches = RecordBatches::try_new(schema1.clone(), vec![batch1.clone()]).unwrap();
1111        let expected = "\
1112+---+-------+
1113| a | b     |
1114+---+-------+
1115| 1 | hello |
1116| 2 | world |
1117+---+-------+";
1118        assert_eq!(batches.pretty_print().unwrap(), expected);
1119
1120        assert_eq!(schema1, batches.schema());
1121        assert_eq!(vec![batch1], batches.take());
1122    }
1123
1124    #[test]
1125    fn test_map_dictionary_to_values_recursively() {
1126        let string_dictionary = ConcreteDataType::dictionary_datatype(
1127            ConcreteDataType::int32_datatype(),
1128            ConcreteDataType::string_datatype(),
1129        );
1130        let list_dictionary =
1131            ConcreteDataType::list_datatype(Arc::new(ConcreteDataType::dictionary_datatype(
1132                ConcreteDataType::uint32_datatype(),
1133                ConcreteDataType::string_datatype(),
1134            )));
1135        let schema = Arc::new(Schema::new(vec![
1136            ColumnSchema::new("host", string_dictionary, true),
1137            ColumnSchema::new("tags", list_dictionary, true),
1138        ]));
1139
1140        let host = DictionaryArray::<Int32Type>::new(
1141            Int32Array::from(vec![0, 1]),
1142            Arc::new(StringArray::from(vec![Some("host-a"), None])),
1143        );
1144        let ArrowDataType::List(item_field) = schema.arrow_schema().field(1).data_type().clone()
1145        else {
1146            unreachable!()
1147        };
1148        let tag_values = DictionaryArray::new(
1149            UInt32Array::from_iter_values([0, 1, 0]),
1150            Arc::new(StringArray::from(vec![Some("a"), None])),
1151        );
1152        let tags = ListArray::new(
1153            item_field,
1154            OffsetBuffer::from_lengths([2, 1]),
1155            Arc::new(tag_values),
1156            None,
1157        );
1158        let batch = DfRecordBatch::try_new(
1159            schema.arrow_schema().clone(),
1160            vec![Arc::new(host), Arc::new(tags)],
1161        )
1162        .unwrap();
1163        let batch = RecordBatch::from_df_record_batch(schema.clone(), batch);
1164
1165        let (mapped_schema, apply_mapper) = map_dictionary_to_values_schema(schema.clone());
1166        assert!(apply_mapper);
1167        assert_eq!(
1168            &ArrowDataType::Utf8,
1169            mapped_schema.arrow_schema().field(0).data_type()
1170        );
1171        assert_eq!(
1172            &ArrowDataType::List(Arc::new(
1173                datatypes::arrow::datatypes::Field::new_list_field(ArrowDataType::Utf8, true,)
1174            )),
1175            mapped_schema.arrow_schema().field(1).data_type()
1176        );
1177
1178        let mapped = map_dictionary_to_values(batch, &schema, &mapped_schema).unwrap();
1179        let host = mapped
1180            .column(0)
1181            .as_any()
1182            .downcast_ref::<StringArray>()
1183            .unwrap();
1184        assert_eq!(vec![Some("host-a"), None], host.iter().collect::<Vec<_>>());
1185        let tags = mapped
1186            .column(1)
1187            .as_any()
1188            .downcast_ref::<ListArray>()
1189            .unwrap();
1190        let tag_values = tags
1191            .values()
1192            .as_any()
1193            .downcast_ref::<StringArray>()
1194            .unwrap();
1195        assert_eq!(
1196            vec![Some("a"), None, Some("a")],
1197            tag_values.iter().collect::<Vec<_>>()
1198        );
1199    }
1200
1201    #[test]
1202    fn test_map_dictionary_to_values_schema_with_duplicate_columns() {
1203        let dictionary_type = ArrowDataType::Dictionary(
1204            Box::new(ArrowDataType::Int32),
1205            Box::new(ArrowDataType::Utf8),
1206        );
1207        let arrow_schema = Arc::new(ArrowSchema::new(vec![
1208            Field::new("area", dictionary_type.clone(), true),
1209            Field::new("area", dictionary_type, true),
1210        ]));
1211        let schema = Arc::new(Schema::try_from(arrow_schema).unwrap());
1212
1213        let (mapped_schema, apply_mapper) = map_dictionary_to_values_schema(schema);
1214        assert!(apply_mapper);
1215        assert_eq!(2, mapped_schema.num_columns());
1216        for field in mapped_schema.arrow_schema().fields() {
1217            assert_eq!("area", field.name());
1218            assert_eq!(&ArrowDataType::Utf8, field.data_type());
1219        }
1220    }
1221
1222    #[tokio::test]
1223    async fn test_simple_recordbatch_stream() {
1224        let column_a = ColumnSchema::new("a", ConcreteDataType::int32_datatype(), false);
1225        let column_b = ColumnSchema::new("b", ConcreteDataType::string_datatype(), false);
1226        let schema = Arc::new(Schema::new(vec![column_a, column_b]));
1227
1228        let va1: VectorRef = Arc::new(Int32Vector::from_slice([1, 2]));
1229        let vb1: VectorRef = Arc::new(StringVector::from(vec!["a", "b"]));
1230        let batch1 = RecordBatch::new(schema.clone(), vec![va1, vb1]).unwrap();
1231
1232        let va2: VectorRef = Arc::new(Int32Vector::from_slice([3, 4, 5]));
1233        let vb2: VectorRef = Arc::new(StringVector::from(vec!["c", "d", "e"]));
1234        let batch2 = RecordBatch::new(schema.clone(), vec![va2, vb2]).unwrap();
1235
1236        let recordbatches =
1237            RecordBatches::try_new(schema.clone(), vec![batch1.clone(), batch2.clone()]).unwrap();
1238        let stream = recordbatches.as_stream();
1239        let collected = util::collect(stream).await.unwrap();
1240        assert_eq!(collected.len(), 2);
1241        assert_eq!(collected[0], batch1);
1242        assert_eq!(collected[1], batch2);
1243    }
1244
1245    const MB: usize = 1024 * 1024;
1246
1247    #[test]
1248    fn test_query_memory_tracker_basic() {
1249        let tracker =
1250            Arc::new(QueryMemoryTracker::builder(10 * MB, OnExhaustedPolicy::Fail).build());
1251
1252        let mut stream1 = tracker.new_stream_tracker();
1253        assert!(stream1.try_track(5 * MB).is_ok());
1254        assert_eq!(tracker.current(), 5 * MB);
1255
1256        let mut stream2 = tracker.new_stream_tracker();
1257        assert!(stream2.try_track(4 * MB).is_ok());
1258        assert_eq!(tracker.current(), 9 * MB);
1259
1260        drop(stream1);
1261        drop(stream2);
1262        assert_eq!(tracker.current(), 0);
1263    }
1264
1265    #[test]
1266    fn test_query_memory_tracker_shared_global_limit() {
1267        let tracker =
1268            Arc::new(QueryMemoryTracker::builder(10 * MB, OnExhaustedPolicy::Fail).build());
1269        let mut stream1 = tracker.new_stream_tracker();
1270        let mut stream2 = tracker.new_stream_tracker();
1271
1272        assert!(stream1.try_track(3 * MB).is_ok());
1273        assert_eq!(tracker.current(), 3 * MB);
1274        assert!(stream2.try_track(6 * MB).is_ok());
1275        assert_eq!(tracker.current(), 9 * MB);
1276
1277        let err = stream2.try_track(2 * MB).unwrap_err();
1278        let err_msg = err.to_string();
1279        assert!(err_msg.contains("6.0MiB used by this stream"));
1280        assert!(err_msg.contains("9.0MiB used globally (90%)"));
1281        assert!(err_msg.contains("hard limit: 10.0MiB"));
1282        assert_eq!(tracker.current(), 9 * MB);
1283
1284        drop(stream1);
1285        assert_eq!(tracker.current(), 6 * MB);
1286        drop(stream2);
1287        assert_eq!(tracker.current(), 0);
1288    }
1289
1290    #[test]
1291    fn test_query_memory_tracker_hard_limit() {
1292        let tracker =
1293            Arc::new(QueryMemoryTracker::builder(10 * MB, OnExhaustedPolicy::Fail).build());
1294        let mut stream = tracker.new_stream_tracker();
1295
1296        assert!(stream.try_track(9 * MB).is_ok());
1297        assert_eq!(tracker.current(), 9 * MB);
1298
1299        assert!(stream.try_track(2 * MB).is_err());
1300        assert_eq!(tracker.current(), 9 * MB);
1301
1302        assert!(stream.try_track(MB).is_ok());
1303        assert_eq!(tracker.current(), 10 * MB);
1304
1305        assert!(stream.try_track(MB).is_err());
1306        assert_eq!(tracker.current(), 10 * MB);
1307
1308        drop(stream);
1309        assert_eq!(tracker.current(), 0);
1310    }
1311
1312    #[test]
1313    fn test_query_memory_tracker_unlimited() {
1314        let tracker = Arc::new(QueryMemoryTracker::builder(0, OnExhaustedPolicy::Fail).build());
1315        let mut stream = tracker.new_stream_tracker();
1316
1317        assert!(stream.try_track(10 * MB).is_ok());
1318        assert_eq!(tracker.current(), 10 * MB);
1319        drop(stream);
1320        assert_eq!(tracker.current(), 0);
1321    }
1322
1323    #[test]
1324    fn test_query_memory_tracker_rounds_to_kilobytes() {
1325        let tracker =
1326            Arc::new(QueryMemoryTracker::builder(10 * MB, OnExhaustedPolicy::Fail).build());
1327        let mut stream = tracker.new_stream_tracker();
1328
1329        assert!(stream.try_track(1_537).is_ok());
1330        assert_eq!(tracker.current(), 2 * 1024);
1331
1332        drop(stream);
1333        assert_eq!(tracker.current(), 0);
1334    }
1335
1336    #[tokio::test]
1337    async fn test_memory_tracked_stream_waits_for_capacity() {
1338        let exhausted = Arc::new(AtomicUsize::new(0));
1339        let rejected = Arc::new(AtomicUsize::new(0));
1340        let exhausted_counter = exhausted.clone();
1341        let rejected_counter = rejected.clone();
1342        let tracker = QueryMemoryTracker::builder(
1343            MB,
1344            OnExhaustedPolicy::Wait {
1345                timeout: Duration::from_millis(200),
1346            },
1347        )
1348        .on_exhausted(move || {
1349            exhausted_counter.fetch_add(1, Ordering::Relaxed);
1350        })
1351        .on_reject(move || {
1352            rejected_counter.fetch_add(1, Ordering::Relaxed);
1353        })
1354        .build();
1355        let batch = large_string_batch(700 * 1024);
1356        let expected_bytes = aligned_tracked_bytes(batch.logical_slice_memory_size());
1357
1358        let mut stream1 = MemoryTrackedStream::new(
1359            RecordBatches::try_new(batch.schema.clone(), vec![batch.clone()])
1360                .unwrap()
1361                .as_stream(),
1362            tracker.clone(),
1363        );
1364        let first = stream1.next().await.unwrap().unwrap();
1365        assert_eq!(first.num_rows(), 1);
1366        assert_eq!(tracker.current(), expected_bytes);
1367
1368        let stream2 = MemoryTrackedStream::new(
1369            RecordBatches::try_new(batch.schema.clone(), vec![batch])
1370                .unwrap()
1371                .as_stream(),
1372            tracker.clone(),
1373        );
1374        let waiter = tokio::spawn(async move {
1375            let mut stream2 = stream2;
1376            stream2.next().await.unwrap()
1377        });
1378
1379        sleep(Duration::from_millis(50)).await;
1380        assert!(!waiter.is_finished());
1381
1382        drop(stream1);
1383        let second = waiter.await.unwrap().unwrap();
1384        assert_eq!(second.num_rows(), 1);
1385        assert_eq!(exhausted.load(Ordering::Relaxed), 1);
1386        assert_eq!(rejected.load(Ordering::Relaxed), 0);
1387    }
1388
1389    #[tokio::test]
1390    async fn test_memory_tracked_stream_wait_times_out() {
1391        let exhausted = Arc::new(AtomicUsize::new(0));
1392        let rejected = Arc::new(AtomicUsize::new(0));
1393        let exhausted_counter = exhausted.clone();
1394        let rejected_counter = rejected.clone();
1395        let tracker = QueryMemoryTracker::builder(
1396            MB,
1397            OnExhaustedPolicy::Wait {
1398                timeout: Duration::from_millis(50),
1399            },
1400        )
1401        .on_exhausted(move || {
1402            exhausted_counter.fetch_add(1, Ordering::Relaxed);
1403        })
1404        .on_reject(move || {
1405            rejected_counter.fetch_add(1, Ordering::Relaxed);
1406        })
1407        .build();
1408        let batch = large_string_batch(700 * 1024);
1409
1410        let mut stream1 = MemoryTrackedStream::new(
1411            RecordBatches::try_new(batch.schema.clone(), vec![batch.clone()])
1412                .unwrap()
1413                .as_stream(),
1414            tracker.clone(),
1415        );
1416        let first = stream1.next().await.unwrap().unwrap();
1417        assert_eq!(first.num_rows(), 1);
1418
1419        let mut stream2 = MemoryTrackedStream::new(
1420            RecordBatches::try_new(batch.schema.clone(), vec![batch])
1421                .unwrap()
1422                .as_stream(),
1423            tracker,
1424        );
1425        let result = timeout(Duration::from_secs(1), stream2.next())
1426            .await
1427            .unwrap();
1428        let error = result.unwrap().unwrap_err();
1429        assert!(error.to_string().contains("timed out waiting"));
1430        assert_eq!(exhausted.load(Ordering::Relaxed), 1);
1431        assert_eq!(rejected.load(Ordering::Relaxed), 1);
1432    }
1433
1434    #[tokio::test]
1435    async fn test_memory_tracked_stream_fail_policy_rejects_immediately() {
1436        let exhausted = Arc::new(AtomicUsize::new(0));
1437        let rejected = Arc::new(AtomicUsize::new(0));
1438        let exhausted_counter = exhausted.clone();
1439        let rejected_counter = rejected.clone();
1440        let tracker = QueryMemoryTracker::builder(MB, OnExhaustedPolicy::Fail)
1441            .on_exhausted(move || {
1442                exhausted_counter.fetch_add(1, Ordering::Relaxed);
1443            })
1444            .on_reject(move || {
1445                rejected_counter.fetch_add(1, Ordering::Relaxed);
1446            })
1447            .build();
1448        let batch = large_string_batch(700 * 1024);
1449
1450        let mut stream1 = MemoryTrackedStream::new(
1451            RecordBatches::try_new(batch.schema.clone(), vec![batch.clone()])
1452                .unwrap()
1453                .as_stream(),
1454            tracker.clone(),
1455        );
1456        let first = stream1.next().await.unwrap().unwrap();
1457        assert_eq!(first.num_rows(), 1);
1458
1459        let mut stream2 = MemoryTrackedStream::new(
1460            RecordBatches::try_new(batch.schema.clone(), vec![batch])
1461                .unwrap()
1462                .as_stream(),
1463            tracker,
1464        );
1465        let result = stream2.next().await.unwrap();
1466        assert!(result.is_err());
1467        assert_eq!(exhausted.load(Ordering::Relaxed), 1);
1468        assert_eq!(rejected.load(Ordering::Relaxed), 1);
1469    }
1470}