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
15pub mod adapter;
16pub mod cursor;
17pub mod error;
18pub mod ext;
19pub mod filter;
20pub mod recordbatch;
21pub mod util;
22
23use std::fmt::{self, Write};
24use std::future::Future;
25use std::pin::Pin;
26use std::sync::Arc;
27
28use adapter::RecordBatchMetrics;
29use arc_swap::ArcSwapOption;
30use common_base::readable_size::ReadableSize;
31use common_error::ext::BoxedError;
32use common_memory_manager::{
33    MemoryGuard, MemoryManager, MemoryMetrics, OnExhaustedPolicy, PermitGranularity,
34};
35use common_telemetry::tracing::Span;
36pub use datafusion::physical_plan::SendableRecordBatchStream as DfSendableRecordBatchStream;
37use datatypes::arrow::array::{Array, ArrayRef, AsArray, StringBuilder};
38use datatypes::arrow::compute::SortOptions;
39use datatypes::arrow::datatypes::{DataType as ArrowDataType, Field};
40use datatypes::arrow::error::ArrowError;
41pub use datatypes::arrow::record_batch::RecordBatch as DfRecordBatch;
42use datatypes::arrow::util::display::{
43    ArrayFormatter, ArrayFormatterFactory, DisplayIndex, FormatOptions, FormatResult,
44};
45use datatypes::arrow::util::pretty::{
46    pretty_format_batches_with_options, pretty_format_batches_with_schema,
47};
48use datatypes::extension::json::is_any_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 result: String = if df_batches.is_empty() {
400            pretty_format_batches_with_schema(self.schema.arrow_schema().clone(), df_batches)
401                .context(error::FormatSnafu)?
402                .to_string()
403        } else {
404            let options =
405                FormatOptions::default().with_formatter_factory(Some(&BinaryFormatterFactory));
406            pretty_format_batches_with_options(df_batches, &options)
407                .context(error::FormatSnafu)?
408                .to_string()
409        };
410
411        Ok(result)
412    }
413
414    pub fn try_new(schema: SchemaRef, batches: Vec<RecordBatch>) -> Result<Self> {
415        for batch in &batches {
416            ensure!(
417                batch.schema == schema,
418                error::CreateRecordBatchesSnafu {
419                    reason: format!(
420                        "expect RecordBatch schema equals {:?}, actual: {:?}",
421                        schema, batch.schema
422                    )
423                }
424            )
425        }
426        Ok(Self { schema, batches })
427    }
428
429    pub fn schema(&self) -> SchemaRef {
430        self.schema.clone()
431    }
432
433    pub fn take(self) -> Vec<RecordBatch> {
434        self.batches
435    }
436
437    pub fn as_stream(&self) -> SendableRecordBatchStream {
438        Box::pin(SimpleRecordBatchStream {
439            inner: RecordBatches {
440                schema: self.schema(),
441                batches: self.batches.clone(),
442            },
443            index: 0,
444        })
445    }
446}
447
448#[derive(Debug)]
449struct BinaryFormatterFactory;
450
451impl ArrayFormatterFactory for BinaryFormatterFactory {
452    fn create_array_formatter<'a>(
453        &self,
454        array: &'a dyn Array,
455        options: &FormatOptions<'a>,
456        field: Option<&'a Field>,
457    ) -> std::result::Result<Option<ArrayFormatter<'a>>, ArrowError> {
458        if !array.data_type().is_binary() {
459            return Ok(None);
460        }
461
462        Ok(Some(ArrayFormatter::new(
463            Box::new(BinaryFormatter {
464                array,
465                is_json: field.is_some_and(is_any_json_extension_type),
466                default: ArrayFormatter::try_new(array, options)?,
467                null: options.null(),
468            }),
469            options.safe(),
470        )))
471    }
472}
473
474struct BinaryFormatter<'a> {
475    array: &'a dyn Array,
476    is_json: bool,
477    default: ArrayFormatter<'a>,
478    null: &'a str,
479}
480
481impl DisplayIndex for BinaryFormatter<'_> {
482    fn write(&self, idx: usize, f: &mut dyn Write) -> FormatResult {
483        if !self.is_json {
484            self.default.value(idx).write(f)?;
485            return Ok(());
486        }
487
488        if self.array.is_null(idx) {
489            write!(f, "{}", self.null)?;
490        } else {
491            let bytes = match self.array.data_type() {
492                ArrowDataType::Binary => self.array.as_binary::<i32>().value(idx),
493                ArrowDataType::LargeBinary => self.array.as_binary::<i64>().value(idx),
494                ArrowDataType::BinaryView => self.array.as_binary_view().value(idx),
495                _ => return Ok(self.default.value(idx).write(f)?),
496            };
497            let value =
498                jsonb_to_string(bytes).map_err(|e| ArrowError::ExternalError(Box::new(e)))?;
499            write!(f, "{value}")?;
500        }
501        Ok(())
502    }
503}
504
505impl IntoIterator for RecordBatches {
506    type Item = RecordBatch;
507    type IntoIter = std::vec::IntoIter<Self::Item>;
508
509    fn into_iter(self) -> Self::IntoIter {
510        self.batches.into_iter()
511    }
512}
513
514pub struct SimpleRecordBatchStream {
515    inner: RecordBatches,
516    index: usize,
517}
518
519impl RecordBatchStream for SimpleRecordBatchStream {
520    fn schema(&self) -> SchemaRef {
521        self.inner.schema()
522    }
523
524    fn output_ordering(&self) -> Option<&[OrderOption]> {
525        None
526    }
527
528    fn metrics(&self) -> Option<RecordBatchMetrics> {
529        None
530    }
531}
532
533impl Stream for SimpleRecordBatchStream {
534    type Item = Result<RecordBatch>;
535
536    fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
537        Poll::Ready(if self.index < self.inner.batches.len() {
538            let batch = self.inner.batches[self.index].clone();
539            self.index += 1;
540            Some(Ok(batch))
541        } else {
542            None
543        })
544    }
545}
546
547/// Adapt a [Stream] of [RecordBatch] to a [RecordBatchStream].
548pub struct RecordBatchStreamWrapper<S> {
549    pub schema: SchemaRef,
550    pub stream: S,
551    pub output_ordering: Option<Vec<OrderOption>>,
552    pub metrics: Arc<ArcSwapOption<RecordBatchMetrics>>,
553    pub span: Span,
554}
555
556impl<S> RecordBatchStreamWrapper<S> {
557    /// Creates a [RecordBatchStreamWrapper] without output ordering requirement.
558    pub fn new(schema: SchemaRef, stream: S) -> RecordBatchStreamWrapper<S> {
559        RecordBatchStreamWrapper {
560            schema,
561            stream,
562            output_ordering: None,
563            metrics: Default::default(),
564            span: Span::current(),
565        }
566    }
567}
568
569impl<S: Stream<Item = Result<RecordBatch>> + Unpin> RecordBatchStream
570    for RecordBatchStreamWrapper<S>
571{
572    fn name(&self) -> &str {
573        "RecordBatchStreamWrapper"
574    }
575
576    fn schema(&self) -> SchemaRef {
577        self.schema.clone()
578    }
579
580    fn output_ordering(&self) -> Option<&[OrderOption]> {
581        self.output_ordering.as_deref()
582    }
583
584    fn metrics(&self) -> Option<RecordBatchMetrics> {
585        self.metrics.load().as_ref().map(|s| s.as_ref().clone())
586    }
587}
588
589impl<S: Stream<Item = Result<RecordBatch>> + Unpin> Stream for RecordBatchStreamWrapper<S> {
590    type Item = Result<RecordBatch>;
591
592    fn poll_next(mut self: Pin<&mut Self>, ctx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
593        let _entered = self.span.clone().entered();
594        Pin::new(&mut self.stream).poll_next(ctx)
595    }
596}
597
598/// Memory tracker for RecordBatch streams. Clone to share the same limit across queries.
599///
600/// Each stream acquires quota independently from this tracker.
601#[derive(Clone)]
602pub struct QueryMemoryTracker {
603    manager: MemoryManager<CallbackMemoryMetrics>,
604    metrics: CallbackMemoryMetrics,
605    on_exhausted_policy: OnExhaustedPolicy,
606}
607
608impl fmt::Debug for QueryMemoryTracker {
609    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
610        f.debug_struct("QueryMemoryTracker")
611            .field("current", &self.current())
612            .field("limit", &self.limit())
613            .field("on_exhausted_policy", &self.on_exhausted_policy)
614            .field("on_update", &self.metrics.has_on_update())
615            .field("on_exhausted", &self.metrics.has_on_exhausted())
616            .field("on_rejected", &self.metrics.has_on_rejected())
617            .finish()
618    }
619}
620
621impl QueryMemoryTracker {
622    /// Create a builder for a query memory tracker.
623    pub fn builder(
624        limit: usize,
625        on_exhausted_policy: OnExhaustedPolicy,
626    ) -> QueryMemoryTrackerBuilder {
627        QueryMemoryTrackerBuilder {
628            limit,
629            on_exhausted_policy,
630            on_update: None,
631            on_exhausted: None,
632            on_reject: None,
633        }
634    }
635
636    fn new_stream_tracker(&self) -> StreamMemoryTracker {
637        StreamMemoryTracker {
638            tracker: self.clone(),
639            guard: self.manager.try_acquire(0).unwrap(),
640            tracked_bytes: 0,
641        }
642    }
643    /// Get the current memory usage in bytes.
644    pub fn current(&self) -> usize {
645        self.manager.used_bytes() as usize
646    }
647
648    fn limit(&self) -> usize {
649        self.manager.limit_bytes() as usize
650    }
651
652    fn reject_error(
653        &self,
654        current: usize,
655        additional: usize,
656        stream_tracked: usize,
657    ) -> error::Error {
658        let limit = self.limit();
659        let msg = format!(
660            "{} requested, {} used globally ({}%), {} used by this stream, hard limit: {}",
661            ReadableSize(additional as u64),
662            ReadableSize(current as u64),
663            (current * 100).checked_div(limit).unwrap_or(0),
664            ReadableSize(stream_tracked as u64),
665            ReadableSize(limit as u64)
666        );
667        error::ExceedMemoryLimitSnafu { msg }.build()
668    }
669
670    fn inc_rejected(&self) {
671        self.metrics.inc_rejected();
672    }
673}
674
675/// Builder for constructing a [`QueryMemoryTracker`] with optional callbacks.
676pub struct QueryMemoryTrackerBuilder {
677    limit: usize,
678    on_exhausted_policy: OnExhaustedPolicy,
679    on_update: Option<UpdateCallback>,
680    on_exhausted: Option<UnitCallback>,
681    on_reject: Option<RejectCallback>,
682}
683
684impl QueryMemoryTrackerBuilder {
685    /// Set a callback to be called whenever the usage changes successfully.
686    /// The callback receives the new total usage in bytes.
687    ///
688    /// # Note
689    /// The callback is called after both successful `track()` and stream drop.
690    /// Usage is exact in unlimited mode and 1KB-aligned in limited mode.
691    pub fn on_update<F>(mut self, on_update: F) -> Self
692    where
693        F: Fn(usize) + Send + Sync + 'static,
694    {
695        self.on_update = Some(Arc::new(on_update));
696        self
697    }
698
699    /// Set a callback to be called when memory is unavailable for immediate acquisition.
700    ///
701    /// # Note
702    /// This is called when the non-blocking allocation fast path fails.
703    /// Requests using `OnExhaustedPolicy::Wait` may still succeed after waiting.
704    /// It is never called when `limit == 0` (unlimited mode).
705    pub fn on_exhausted<F>(mut self, on_exhausted: F) -> Self
706    where
707        F: Fn() + Send + Sync + 'static,
708    {
709        self.on_exhausted = Some(Arc::new(on_exhausted));
710        self
711    }
712
713    /// Set a callback to be called when the request ultimately fails due to memory pressure.
714    pub fn on_reject<F>(mut self, on_reject: F) -> Self
715    where
716        F: Fn() + Send + Sync + 'static,
717    {
718        self.on_reject = Some(Arc::new(on_reject));
719        self
720    }
721
722    /// Build a [`QueryMemoryTracker`] from this builder.
723    pub fn build(self) -> QueryMemoryTracker {
724        let metrics = CallbackMemoryMetrics::new(self.on_update, self.on_exhausted, self.on_reject);
725        let manager = MemoryManager::with_granularity(
726            self.limit as u64,
727            PermitGranularity::Kilobyte,
728            metrics.clone(),
729        );
730
731        QueryMemoryTracker {
732            manager,
733            metrics,
734            on_exhausted_policy: self.on_exhausted_policy,
735        }
736    }
737}
738
739struct StreamMemoryTracker {
740    tracker: QueryMemoryTracker,
741    guard: MemoryGuard<CallbackMemoryMetrics>,
742    tracked_bytes: usize,
743}
744
745type MemoryAcquireResult = std::result::Result<(), common_memory_manager::Error>;
746
747impl StreamMemoryTracker {
748    fn inc_rejected(&self) {
749        self.tracker.inc_rejected();
750    }
751
752    fn try_track(&mut self, additional: usize) -> Result<()> {
753        if self.guard.try_acquire_additional(additional as u64) {
754            self.tracked_bytes = self.tracked_bytes.saturating_add(additional);
755            Ok(())
756        } else {
757            Err(self.reject_error(additional))
758        }
759    }
760
761    async fn track_with_policy(mut self, additional: usize) -> (Self, MemoryAcquireResult) {
762        let result = self
763            .guard
764            .acquire_additional_with_policy(additional as u64, self.tracker.on_exhausted_policy)
765            .await;
766        if result.is_ok() {
767            self.tracked_bytes = self.tracked_bytes.saturating_add(additional);
768        }
769        (self, result)
770    }
771
772    fn reject_error(&self, additional: usize) -> error::Error {
773        let current = self.tracker.current();
774        self.tracker
775            .reject_error(current, additional, self.tracked_bytes)
776    }
777
778    fn wait_error(&self, additional: usize, source: common_memory_manager::Error) -> error::Error {
779        match source {
780            common_memory_manager::Error::MemoryLimitExceeded { .. } => {
781                self.reject_error(additional)
782            }
783            common_memory_manager::Error::MemoryAcquireTimeout { waited, .. } => {
784                let current = self.tracker.current();
785                let limit = self.tracker.limit();
786                let msg = format!(
787                    "timed out waiting {:?} for {}, {} used globally ({}%), {} used by this stream, hard limit: {}",
788                    waited,
789                    ReadableSize(additional as u64),
790                    ReadableSize(current as u64),
791                    (current * 100).checked_div(limit).unwrap_or(0),
792                    ReadableSize(self.tracked_bytes as u64),
793                    ReadableSize(limit as u64)
794                );
795                error::ExceedMemoryLimitSnafu { msg }.build()
796            }
797            error => error::ExternalSnafu.into_error(BoxedError::new(error)),
798        }
799    }
800}
801
802type PendingTrackFuture = Pin<
803    Box<dyn Future<Output = (StreamMemoryTracker, RecordBatch, usize, MemoryAcquireResult)> + Send>,
804>;
805
806#[derive(Clone)]
807struct CallbackMemoryMetrics {
808    inner: Arc<CallbackMemoryMetricsInner>,
809}
810
811type UpdateCallback = Arc<dyn Fn(usize) + Send + Sync>;
812type UnitCallback = Arc<dyn Fn() + Send + Sync>;
813type RejectCallback = UnitCallback;
814
815struct CallbackMemoryMetricsInner {
816    on_update: Option<UpdateCallback>,
817    on_exhausted: Option<UnitCallback>,
818    on_reject: Option<RejectCallback>,
819}
820
821impl CallbackMemoryMetrics {
822    fn new(
823        on_update: Option<UpdateCallback>,
824        on_exhausted: Option<UnitCallback>,
825        on_reject: Option<RejectCallback>,
826    ) -> Self {
827        Self {
828            inner: Arc::new(CallbackMemoryMetricsInner {
829                on_update,
830                on_exhausted,
831                on_reject,
832            }),
833        }
834    }
835
836    fn has_on_update(&self) -> bool {
837        self.inner.on_update.is_some()
838    }
839
840    fn has_on_exhausted(&self) -> bool {
841        self.inner.on_exhausted.is_some()
842    }
843
844    fn has_on_rejected(&self) -> bool {
845        self.inner.on_reject.is_some()
846    }
847
848    fn inc_rejected(&self) {
849        if let Some(callback) = &self.inner.on_reject {
850            callback();
851        }
852    }
853}
854
855impl MemoryMetrics for CallbackMemoryMetrics {
856    fn set_limit(&self, _: i64) {}
857
858    fn set_in_use(&self, bytes: i64) {
859        if let Some(callback) = &self.inner.on_update {
860            callback(bytes.max(0) as usize);
861        }
862    }
863
864    fn inc_exhausted(&self, _: &str) {
865        if let Some(callback) = &self.inner.on_exhausted {
866            callback();
867        }
868    }
869}
870
871/// A wrapper stream that tracks memory usage of RecordBatches.
872pub struct MemoryTrackedStream {
873    inner: SendableRecordBatchStream,
874    tracker: Option<StreamMemoryTracker>,
875    // Waiting stores a batch that has already been pulled from the inner stream but has not yet
876    // acquired additional quota. This keeps `poll_next()` non-blocking and allows bounded waits,
877    // at the cost of temporarily holding one untracked batch per blocked stream in memory.
878    waiting: Option<PendingTrackFuture>,
879}
880
881impl MemoryTrackedStream {
882    pub fn new(inner: SendableRecordBatchStream, tracker: QueryMemoryTracker) -> Self {
883        Self {
884            inner,
885            tracker: Some(tracker.new_stream_tracker()),
886            waiting: None,
887        }
888    }
889
890    fn ready_tracker_mut(&mut self) -> &mut StreamMemoryTracker {
891        debug_assert!(
892            self.waiting.is_none(),
893            "a ready tracker must not coexist with a waiting future"
894        );
895        self.tracker.as_mut().unwrap()
896    }
897
898    fn enter_waiting(&mut self, batch: RecordBatch, additional: usize) {
899        debug_assert!(
900            self.waiting.is_none(),
901            "enter_waiting should only be called from the ready state"
902        );
903        debug_assert!(
904            self.tracker.is_some(),
905            "enter_waiting requires a tracker in the ready state"
906        );
907        let tracker = self.tracker.take().unwrap();
908        self.waiting = Some(Self::start_waiting(tracker, batch, additional));
909    }
910
911    fn start_waiting(
912        tracker: StreamMemoryTracker,
913        batch: RecordBatch,
914        additional: usize,
915    ) -> PendingTrackFuture {
916        Box::pin(async move {
917            let (tracker, result) = tracker.track_with_policy(additional).await;
918            (tracker, batch, additional, result)
919        })
920    }
921
922    fn poll_waiting(&mut self, cx: &mut Context<'_>) -> Poll<Option<Result<RecordBatch>>> {
923        let future = self.waiting.as_mut().unwrap();
924        match future.as_mut().poll(cx) {
925            Poll::Ready((tracker, batch, additional, result)) => {
926                let output = match result {
927                    Ok(()) => Ok(batch),
928                    Err(error) => {
929                        tracker.inc_rejected();
930                        Err(tracker.wait_error(additional, error))
931                    }
932                };
933                self.waiting = None;
934                self.tracker = Some(tracker);
935                Poll::Ready(Some(output))
936            }
937            Poll::Pending => Poll::Pending,
938        }
939    }
940
941    fn poll_batch(
942        &mut self,
943        batch: RecordBatch,
944        cx: &mut Context<'_>,
945    ) -> Poll<Option<Result<RecordBatch>>> {
946        let additional = batch.logical_slice_memory_size();
947        let tracker = self.ready_tracker_mut();
948
949        if let Err(error) = tracker.try_track(additional) {
950            match tracker.tracker.on_exhausted_policy {
951                OnExhaustedPolicy::Fail => {
952                    tracker.inc_rejected();
953                    return Poll::Ready(Some(Err(error)));
954                }
955                // `Wait` is a deliberate tradeoff: the batch has already been materialized, so we
956                // keep it in memory while waiting for quota instead of failing immediately. Under
957                // contention, real memory usage can therefore exceed `scan_memory_limit` by up to
958                // one buffered batch per blocked stream.
959                OnExhaustedPolicy::Wait { .. } => {
960                    self.enter_waiting(batch, additional);
961                    return self.poll_waiting(cx);
962                }
963            }
964        }
965
966        Poll::Ready(Some(Ok(batch)))
967    }
968}
969
970impl Stream for MemoryTrackedStream {
971    type Item = Result<RecordBatch>;
972
973    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
974        if self.waiting.is_some() {
975            return self.poll_waiting(cx);
976        }
977
978        match Pin::new(&mut self.inner).poll_next(cx) {
979            Poll::Ready(Some(Ok(batch))) => self.poll_batch(batch, cx),
980            Poll::Ready(Some(Err(error))) => Poll::Ready(Some(Err(error))),
981            Poll::Ready(None) => Poll::Ready(None),
982            Poll::Pending => Poll::Pending,
983        }
984    }
985
986    fn size_hint(&self) -> (usize, Option<usize>) {
987        self.inner.size_hint()
988    }
989}
990
991impl RecordBatchStream for MemoryTrackedStream {
992    fn schema(&self) -> SchemaRef {
993        self.inner.schema()
994    }
995
996    fn output_ordering(&self) -> Option<&[OrderOption]> {
997        self.inner.output_ordering()
998    }
999
1000    fn metrics(&self) -> Option<RecordBatchMetrics> {
1001        self.inner.metrics()
1002    }
1003}
1004
1005#[cfg(test)]
1006mod tests {
1007    use std::sync::Arc;
1008    use std::sync::atomic::{AtomicUsize, Ordering};
1009    use std::time::Duration;
1010
1011    use common_memory_manager::{OnExhaustedPolicy, PermitGranularity};
1012    use datatypes::arrow::array::{
1013        DictionaryArray, Int32Array, ListArray, StringArray, UInt32Array,
1014    };
1015    use datatypes::arrow::buffer::OffsetBuffer;
1016    use datatypes::arrow::datatypes::{
1017        DataType as ArrowDataType, Field, Int32Type, Schema as ArrowSchema,
1018    };
1019    use datatypes::prelude::{ConcreteDataType, VectorRef};
1020    use datatypes::schema::{ColumnSchema, Schema};
1021    use datatypes::vectors::{BooleanVector, Int32Vector, StringVector};
1022    use futures::StreamExt;
1023    use tokio::time::{sleep, timeout};
1024
1025    use super::*;
1026
1027    fn large_string_batch(bytes: usize) -> RecordBatch {
1028        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1029            "payload",
1030            ConcreteDataType::string_datatype(),
1031            false,
1032        )]));
1033        let payload = "x".repeat(bytes);
1034        let vector: VectorRef = Arc::new(StringVector::from(vec![payload]));
1035        RecordBatch::new(schema, vec![vector]).unwrap()
1036    }
1037
1038    fn aligned_tracked_bytes(bytes: usize) -> usize {
1039        PermitGranularity::Kilobyte
1040            .permits_to_bytes(PermitGranularity::Kilobyte.bytes_to_permits(bytes as u64))
1041            as usize
1042    }
1043
1044    #[tokio::test]
1045    async fn test_memory_tracked_stream_charges_logical_slice_size() {
1046        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1047            "payload",
1048            ConcreteDataType::string_datatype(),
1049            false,
1050        )]));
1051        let payloads: Vec<_> = (0..1024)
1052            .map(|value| format!("payload-{value:04}"))
1053            .collect();
1054        let batch = RecordBatch::new(schema, vec![Arc::new(StringVector::from(payloads)) as _])
1055            .unwrap()
1056            .slice(512, 1)
1057            .unwrap();
1058        let expected_bytes = aligned_tracked_bytes(batch.logical_slice_memory_size());
1059        assert!(expected_bytes < aligned_tracked_bytes(batch.buffer_memory_size()));
1060        let tracker = QueryMemoryTracker::builder(MB, OnExhaustedPolicy::Fail).build();
1061        let mut stream = MemoryTrackedStream::new(
1062            RecordBatches::try_new(batch.schema.clone(), vec![batch])
1063                .unwrap()
1064                .as_stream(),
1065            tracker.clone(),
1066        );
1067
1068        stream.next().await.unwrap().unwrap();
1069
1070        assert_eq!(tracker.current(), expected_bytes);
1071    }
1072
1073    #[test]
1074    fn test_recordbatches_try_from_columns() {
1075        let schema = Arc::new(Schema::new(vec![ColumnSchema::new(
1076            "a",
1077            ConcreteDataType::int32_datatype(),
1078            false,
1079        )]));
1080        let result = RecordBatches::try_from_columns(
1081            schema.clone(),
1082            vec![Arc::new(StringVector::from(vec!["hello", "world"])) as _],
1083        );
1084        assert!(result.is_err());
1085
1086        let v: VectorRef = Arc::new(Int32Vector::from_slice([1, 2]));
1087        let expected = vec![RecordBatch::new(schema.clone(), vec![v.clone()]).unwrap()];
1088        let r = RecordBatches::try_from_columns(schema, vec![v]).unwrap();
1089        assert_eq!(r.take(), expected);
1090    }
1091
1092    #[tokio::test]
1093    async fn test_recordbatches_pretty_print_empty_batches_preserves_schema() {
1094        let schema = Arc::new(Schema::new(vec![
1095            ColumnSchema::new("unit", ConcreteDataType::string_datatype(), false),
1096            ColumnSchema::new(
1097                "ts",
1098                ConcreteDataType::timestamp_millisecond_datatype(),
1099                false,
1100            ),
1101            ColumnSchema::new(
1102                "lhs.degrees(val) + rhs.radians(val)",
1103                ConcreteDataType::float64_datatype(),
1104                false,
1105            ),
1106        ]));
1107        let batches =
1108            RecordBatches::try_collect(Box::pin(EmptyRecordBatchStream::new(schema.clone())))
1109                .await
1110                .unwrap();
1111
1112        assert_eq!(schema, batches.schema());
1113        let expected = "\
1114+------+----+-------------------------------------+
1115| unit | ts | lhs.degrees(val) + rhs.radians(val) |
1116+------+----+-------------------------------------+
1117+------+----+-------------------------------------+";
1118        assert_eq!(expected, batches.pretty_print().unwrap());
1119    }
1120
1121    #[test]
1122    fn test_recordbatches_try_new() {
1123        let column_a = ColumnSchema::new("a", ConcreteDataType::int32_datatype(), false);
1124        let column_b = ColumnSchema::new("b", ConcreteDataType::string_datatype(), false);
1125        let column_c = ColumnSchema::new("c", ConcreteDataType::boolean_datatype(), false);
1126
1127        let va: VectorRef = Arc::new(Int32Vector::from_slice([1, 2]));
1128        let vb: VectorRef = Arc::new(StringVector::from(vec!["hello", "world"]));
1129        let vc: VectorRef = Arc::new(BooleanVector::from(vec![true, false]));
1130
1131        let schema1 = Arc::new(Schema::new(vec![column_a.clone(), column_b]));
1132        let batch1 = RecordBatch::new(schema1.clone(), vec![va.clone(), vb]).unwrap();
1133
1134        let schema2 = Arc::new(Schema::new(vec![column_a, column_c]));
1135        let batch2 = RecordBatch::new(schema2.clone(), vec![va, vc]).unwrap();
1136
1137        let result = RecordBatches::try_new(schema1.clone(), vec![batch1.clone(), batch2]);
1138        assert!(result.is_err());
1139        assert_eq!(
1140            result.unwrap_err().to_string(),
1141            format!(
1142                "Failed to create RecordBatches, reason: expect RecordBatch schema equals {schema1:?}, actual: {schema2:?}",
1143            )
1144        );
1145
1146        let batches = RecordBatches::try_new(schema1.clone(), vec![batch1.clone()]).unwrap();
1147        let expected = "\
1148+---+-------+
1149| a | b     |
1150+---+-------+
1151| 1 | hello |
1152| 2 | world |
1153+---+-------+";
1154        assert_eq!(batches.pretty_print().unwrap(), expected);
1155
1156        assert_eq!(schema1, batches.schema());
1157        assert_eq!(vec![batch1], batches.take());
1158    }
1159
1160    #[test]
1161    fn test_map_dictionary_to_values_recursively() {
1162        let string_dictionary = ConcreteDataType::dictionary_datatype(
1163            ConcreteDataType::int32_datatype(),
1164            ConcreteDataType::string_datatype(),
1165        );
1166        let list_dictionary =
1167            ConcreteDataType::list_datatype(Arc::new(ConcreteDataType::dictionary_datatype(
1168                ConcreteDataType::uint32_datatype(),
1169                ConcreteDataType::string_datatype(),
1170            )));
1171        let schema = Arc::new(Schema::new(vec![
1172            ColumnSchema::new("host", string_dictionary, true),
1173            ColumnSchema::new("tags", list_dictionary, true),
1174        ]));
1175
1176        let host = DictionaryArray::<Int32Type>::new(
1177            Int32Array::from(vec![0, 1]),
1178            Arc::new(StringArray::from(vec![Some("host-a"), None])),
1179        );
1180        let ArrowDataType::List(item_field) = schema.arrow_schema().field(1).data_type().clone()
1181        else {
1182            unreachable!()
1183        };
1184        let tag_values = DictionaryArray::new(
1185            UInt32Array::from_iter_values([0, 1, 0]),
1186            Arc::new(StringArray::from(vec![Some("a"), None])),
1187        );
1188        let tags = ListArray::new(
1189            item_field,
1190            OffsetBuffer::from_lengths([2, 1]),
1191            Arc::new(tag_values),
1192            None,
1193        );
1194        let batch = DfRecordBatch::try_new(
1195            schema.arrow_schema().clone(),
1196            vec![Arc::new(host), Arc::new(tags)],
1197        )
1198        .unwrap();
1199        let batch = RecordBatch::from_df_record_batch(schema.clone(), batch);
1200
1201        let (mapped_schema, apply_mapper) = map_dictionary_to_values_schema(schema.clone());
1202        assert!(apply_mapper);
1203        assert_eq!(
1204            &ArrowDataType::Utf8,
1205            mapped_schema.arrow_schema().field(0).data_type()
1206        );
1207        assert_eq!(
1208            &ArrowDataType::List(Arc::new(
1209                datatypes::arrow::datatypes::Field::new_list_field(ArrowDataType::Utf8, true,)
1210            )),
1211            mapped_schema.arrow_schema().field(1).data_type()
1212        );
1213
1214        let mapped = map_dictionary_to_values(batch, &schema, &mapped_schema).unwrap();
1215        let host = mapped
1216            .column(0)
1217            .as_any()
1218            .downcast_ref::<StringArray>()
1219            .unwrap();
1220        assert_eq!(vec![Some("host-a"), None], host.iter().collect::<Vec<_>>());
1221        let tags = mapped
1222            .column(1)
1223            .as_any()
1224            .downcast_ref::<ListArray>()
1225            .unwrap();
1226        let tag_values = tags
1227            .values()
1228            .as_any()
1229            .downcast_ref::<StringArray>()
1230            .unwrap();
1231        assert_eq!(
1232            vec![Some("a"), None, Some("a")],
1233            tag_values.iter().collect::<Vec<_>>()
1234        );
1235    }
1236
1237    #[test]
1238    fn test_map_dictionary_to_values_schema_with_duplicate_columns() {
1239        let dictionary_type = ArrowDataType::Dictionary(
1240            Box::new(ArrowDataType::Int32),
1241            Box::new(ArrowDataType::Utf8),
1242        );
1243        let arrow_schema = Arc::new(ArrowSchema::new(vec![
1244            Field::new("area", dictionary_type.clone(), true),
1245            Field::new("area", dictionary_type, true),
1246        ]));
1247        let schema = Arc::new(Schema::try_from(arrow_schema).unwrap());
1248
1249        let (mapped_schema, apply_mapper) = map_dictionary_to_values_schema(schema);
1250        assert!(apply_mapper);
1251        assert_eq!(2, mapped_schema.num_columns());
1252        for field in mapped_schema.arrow_schema().fields() {
1253            assert_eq!("area", field.name());
1254            assert_eq!(&ArrowDataType::Utf8, field.data_type());
1255        }
1256    }
1257
1258    #[tokio::test]
1259    async fn test_simple_recordbatch_stream() {
1260        let column_a = ColumnSchema::new("a", ConcreteDataType::int32_datatype(), false);
1261        let column_b = ColumnSchema::new("b", ConcreteDataType::string_datatype(), false);
1262        let schema = Arc::new(Schema::new(vec![column_a, column_b]));
1263
1264        let va1: VectorRef = Arc::new(Int32Vector::from_slice([1, 2]));
1265        let vb1: VectorRef = Arc::new(StringVector::from(vec!["a", "b"]));
1266        let batch1 = RecordBatch::new(schema.clone(), vec![va1, vb1]).unwrap();
1267
1268        let va2: VectorRef = Arc::new(Int32Vector::from_slice([3, 4, 5]));
1269        let vb2: VectorRef = Arc::new(StringVector::from(vec!["c", "d", "e"]));
1270        let batch2 = RecordBatch::new(schema.clone(), vec![va2, vb2]).unwrap();
1271
1272        let recordbatches =
1273            RecordBatches::try_new(schema.clone(), vec![batch1.clone(), batch2.clone()]).unwrap();
1274        let stream = recordbatches.as_stream();
1275        let collected = util::collect(stream).await.unwrap();
1276        assert_eq!(collected.len(), 2);
1277        assert_eq!(collected[0], batch1);
1278        assert_eq!(collected[1], batch2);
1279    }
1280
1281    const MB: usize = 1024 * 1024;
1282
1283    #[test]
1284    fn test_query_memory_tracker_basic() {
1285        let tracker =
1286            Arc::new(QueryMemoryTracker::builder(10 * MB, OnExhaustedPolicy::Fail).build());
1287
1288        let mut stream1 = tracker.new_stream_tracker();
1289        assert!(stream1.try_track(5 * MB).is_ok());
1290        assert_eq!(tracker.current(), 5 * MB);
1291
1292        let mut stream2 = tracker.new_stream_tracker();
1293        assert!(stream2.try_track(4 * MB).is_ok());
1294        assert_eq!(tracker.current(), 9 * MB);
1295
1296        drop(stream1);
1297        drop(stream2);
1298        assert_eq!(tracker.current(), 0);
1299    }
1300
1301    #[test]
1302    fn test_query_memory_tracker_shared_global_limit() {
1303        let tracker =
1304            Arc::new(QueryMemoryTracker::builder(10 * MB, OnExhaustedPolicy::Fail).build());
1305        let mut stream1 = tracker.new_stream_tracker();
1306        let mut stream2 = tracker.new_stream_tracker();
1307
1308        assert!(stream1.try_track(3 * MB).is_ok());
1309        assert_eq!(tracker.current(), 3 * MB);
1310        assert!(stream2.try_track(6 * MB).is_ok());
1311        assert_eq!(tracker.current(), 9 * MB);
1312
1313        let err = stream2.try_track(2 * MB).unwrap_err();
1314        let err_msg = err.to_string();
1315        assert!(err_msg.contains("6.0MiB used by this stream"));
1316        assert!(err_msg.contains("9.0MiB used globally (90%)"));
1317        assert!(err_msg.contains("hard limit: 10.0MiB"));
1318        assert_eq!(tracker.current(), 9 * MB);
1319
1320        drop(stream1);
1321        assert_eq!(tracker.current(), 6 * MB);
1322        drop(stream2);
1323        assert_eq!(tracker.current(), 0);
1324    }
1325
1326    #[test]
1327    fn test_query_memory_tracker_hard_limit() {
1328        let tracker =
1329            Arc::new(QueryMemoryTracker::builder(10 * MB, OnExhaustedPolicy::Fail).build());
1330        let mut stream = tracker.new_stream_tracker();
1331
1332        assert!(stream.try_track(9 * MB).is_ok());
1333        assert_eq!(tracker.current(), 9 * MB);
1334
1335        assert!(stream.try_track(2 * MB).is_err());
1336        assert_eq!(tracker.current(), 9 * MB);
1337
1338        assert!(stream.try_track(MB).is_ok());
1339        assert_eq!(tracker.current(), 10 * MB);
1340
1341        assert!(stream.try_track(MB).is_err());
1342        assert_eq!(tracker.current(), 10 * MB);
1343
1344        drop(stream);
1345        assert_eq!(tracker.current(), 0);
1346    }
1347
1348    #[test]
1349    fn test_query_memory_tracker_unlimited() {
1350        let tracker = Arc::new(QueryMemoryTracker::builder(0, OnExhaustedPolicy::Fail).build());
1351        let mut stream = tracker.new_stream_tracker();
1352
1353        assert!(stream.try_track(10 * MB).is_ok());
1354        assert_eq!(tracker.current(), 10 * MB);
1355        drop(stream);
1356        assert_eq!(tracker.current(), 0);
1357    }
1358
1359    #[test]
1360    fn test_query_memory_tracker_rounds_to_kilobytes() {
1361        let tracker =
1362            Arc::new(QueryMemoryTracker::builder(10 * MB, OnExhaustedPolicy::Fail).build());
1363        let mut stream = tracker.new_stream_tracker();
1364
1365        assert!(stream.try_track(1_537).is_ok());
1366        assert_eq!(tracker.current(), 2 * 1024);
1367
1368        drop(stream);
1369        assert_eq!(tracker.current(), 0);
1370    }
1371
1372    #[tokio::test]
1373    async fn test_memory_tracked_stream_waits_for_capacity() {
1374        let exhausted = Arc::new(AtomicUsize::new(0));
1375        let rejected = Arc::new(AtomicUsize::new(0));
1376        let exhausted_counter = exhausted.clone();
1377        let rejected_counter = rejected.clone();
1378        let tracker = QueryMemoryTracker::builder(
1379            MB,
1380            OnExhaustedPolicy::Wait {
1381                timeout: Duration::from_millis(200),
1382            },
1383        )
1384        .on_exhausted(move || {
1385            exhausted_counter.fetch_add(1, Ordering::Relaxed);
1386        })
1387        .on_reject(move || {
1388            rejected_counter.fetch_add(1, Ordering::Relaxed);
1389        })
1390        .build();
1391        let batch = large_string_batch(700 * 1024);
1392        let expected_bytes = aligned_tracked_bytes(batch.logical_slice_memory_size());
1393
1394        let mut stream1 = MemoryTrackedStream::new(
1395            RecordBatches::try_new(batch.schema.clone(), vec![batch.clone()])
1396                .unwrap()
1397                .as_stream(),
1398            tracker.clone(),
1399        );
1400        let first = stream1.next().await.unwrap().unwrap();
1401        assert_eq!(first.num_rows(), 1);
1402        assert_eq!(tracker.current(), expected_bytes);
1403
1404        let stream2 = MemoryTrackedStream::new(
1405            RecordBatches::try_new(batch.schema.clone(), vec![batch])
1406                .unwrap()
1407                .as_stream(),
1408            tracker.clone(),
1409        );
1410        let waiter = tokio::spawn(async move {
1411            let mut stream2 = stream2;
1412            stream2.next().await.unwrap()
1413        });
1414
1415        sleep(Duration::from_millis(50)).await;
1416        assert!(!waiter.is_finished());
1417
1418        drop(stream1);
1419        let second = waiter.await.unwrap().unwrap();
1420        assert_eq!(second.num_rows(), 1);
1421        assert_eq!(exhausted.load(Ordering::Relaxed), 1);
1422        assert_eq!(rejected.load(Ordering::Relaxed), 0);
1423    }
1424
1425    #[tokio::test]
1426    async fn test_memory_tracked_stream_wait_times_out() {
1427        let exhausted = Arc::new(AtomicUsize::new(0));
1428        let rejected = Arc::new(AtomicUsize::new(0));
1429        let exhausted_counter = exhausted.clone();
1430        let rejected_counter = rejected.clone();
1431        let tracker = QueryMemoryTracker::builder(
1432            MB,
1433            OnExhaustedPolicy::Wait {
1434                timeout: Duration::from_millis(50),
1435            },
1436        )
1437        .on_exhausted(move || {
1438            exhausted_counter.fetch_add(1, Ordering::Relaxed);
1439        })
1440        .on_reject(move || {
1441            rejected_counter.fetch_add(1, Ordering::Relaxed);
1442        })
1443        .build();
1444        let batch = large_string_batch(700 * 1024);
1445
1446        let mut stream1 = MemoryTrackedStream::new(
1447            RecordBatches::try_new(batch.schema.clone(), vec![batch.clone()])
1448                .unwrap()
1449                .as_stream(),
1450            tracker.clone(),
1451        );
1452        let first = stream1.next().await.unwrap().unwrap();
1453        assert_eq!(first.num_rows(), 1);
1454
1455        let mut stream2 = MemoryTrackedStream::new(
1456            RecordBatches::try_new(batch.schema.clone(), vec![batch])
1457                .unwrap()
1458                .as_stream(),
1459            tracker,
1460        );
1461        let result = timeout(Duration::from_secs(1), stream2.next())
1462            .await
1463            .unwrap();
1464        let error = result.unwrap().unwrap_err();
1465        assert!(error.to_string().contains("timed out waiting"));
1466        assert_eq!(exhausted.load(Ordering::Relaxed), 1);
1467        assert_eq!(rejected.load(Ordering::Relaxed), 1);
1468    }
1469
1470    #[tokio::test]
1471    async fn test_memory_tracked_stream_fail_policy_rejects_immediately() {
1472        let exhausted = Arc::new(AtomicUsize::new(0));
1473        let rejected = Arc::new(AtomicUsize::new(0));
1474        let exhausted_counter = exhausted.clone();
1475        let rejected_counter = rejected.clone();
1476        let tracker = QueryMemoryTracker::builder(MB, OnExhaustedPolicy::Fail)
1477            .on_exhausted(move || {
1478                exhausted_counter.fetch_add(1, Ordering::Relaxed);
1479            })
1480            .on_reject(move || {
1481                rejected_counter.fetch_add(1, Ordering::Relaxed);
1482            })
1483            .build();
1484        let batch = large_string_batch(700 * 1024);
1485
1486        let mut stream1 = MemoryTrackedStream::new(
1487            RecordBatches::try_new(batch.schema.clone(), vec![batch.clone()])
1488                .unwrap()
1489                .as_stream(),
1490            tracker.clone(),
1491        );
1492        let first = stream1.next().await.unwrap().unwrap();
1493        assert_eq!(first.num_rows(), 1);
1494
1495        let mut stream2 = MemoryTrackedStream::new(
1496            RecordBatches::try_new(batch.schema.clone(), vec![batch])
1497                .unwrap()
1498                .as_stream(),
1499            tracker,
1500        );
1501        let result = stream2.next().await.unwrap();
1502        assert!(result.is_err());
1503        assert_eq!(exhausted.load(Ordering::Relaxed), 1);
1504        assert_eq!(rejected.load(Ordering::Relaxed), 1);
1505    }
1506}