1pub 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
80pub struct SendableRecordBatchMapper {
87 inner: SendableRecordBatchStream,
88 mapper: fn(RecordBatch, &SchemaRef, &SchemaRef) -> Result<RecordBatch>,
91 schema: SchemaRef,
93 apply_mapper: bool,
95}
96
97pub 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
147pub 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
172pub 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
200pub 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 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
244pub 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 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
325pub struct EmptyRecordBatchStream {
328 schema: SchemaRef,
330}
331
332impl EmptyRecordBatchStream {
333 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
547pub 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 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#[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 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 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
675pub 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 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 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 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 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
871pub struct MemoryTrackedStream {
873 inner: SendableRecordBatchStream,
874 tracker: Option<StreamMemoryTracker>,
875 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 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}