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