use std::{
fmt::{Debug, Display, Formatter},
sync::{
Arc,
atomic::{AtomicUsize, Ordering},
},
};
use super::{MemoryConsumer, MemoryLimit, MemoryPool, MemoryReservation};
use datafusion_common::Result;
pub struct PeakRecordingPool {
inner: Arc<dyn MemoryPool>,
reserved: AtomicUsize,
peak: AtomicUsize,
max: AtomicUsize,
}
impl PeakRecordingPool {
pub fn new(inner: Arc<dyn MemoryPool>) -> Self {
Self {
inner,
reserved: AtomicUsize::new(0),
peak: AtomicUsize::new(0),
max: AtomicUsize::new(0),
}
}
pub fn from_pool(pool: &dyn MemoryPool) -> Option<&Self> {
pool.downcast_ref::<Self>()
}
pub fn peak_reserved(&self) -> usize {
self.peak.load(Ordering::Relaxed)
}
pub fn max_reserved(&self) -> usize {
self.max.load(Ordering::Relaxed)
}
pub fn reset_peak(&self) {
self.peak
.store(self.reserved.load(Ordering::Relaxed), Ordering::Relaxed);
}
fn record(&self, additional: usize) {
let reserved =
self.reserved.fetch_add(additional, Ordering::Relaxed) + additional;
self.peak.fetch_max(reserved, Ordering::Relaxed);
self.max.fetch_max(reserved, Ordering::Relaxed);
}
}
impl Debug for PeakRecordingPool {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PeakRecordingPool")
.field("inner", &self.inner)
.field("peak", &self.peak_reserved())
.field("max", &self.max_reserved())
.finish()
}
}
impl Display for PeakRecordingPool {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
Display::fmt(&self.inner, f)
}
}
impl MemoryPool for PeakRecordingPool {
fn name(&self) -> &str {
self.inner.name()
}
fn register(&self, consumer: &MemoryConsumer) {
self.inner.register(consumer);
}
fn unregister(&self, consumer: &MemoryConsumer) {
self.inner.unregister(consumer);
}
fn grow(&self, reservation: &MemoryReservation, additional: usize) {
self.inner.grow(reservation, additional);
self.record(additional);
}
fn shrink(&self, reservation: &MemoryReservation, shrink: usize) {
self.inner.shrink(reservation, shrink);
self.reserved.fetch_sub(shrink, Ordering::Relaxed);
}
fn try_grow(&self, reservation: &MemoryReservation, additional: usize) -> Result<()> {
self.inner.try_grow(reservation, additional)?;
self.record(additional);
Ok(())
}
fn reserved(&self) -> usize {
self.inner.reserved()
}
fn memory_limit(&self) -> MemoryLimit {
self.inner.memory_limit()
}
}
#[cfg(test)]
mod tests {
use crate::memory_pool::GreedyMemoryPool;
use super::*;
fn pool(limit: usize) -> (Arc<PeakRecordingPool>, Arc<dyn MemoryPool>) {
let recording = Arc::new(PeakRecordingPool::new(Arc::new(
GreedyMemoryPool::new(limit),
)));
let pool = Arc::clone(&recording) as Arc<dyn MemoryPool>;
(recording, pool)
}
#[test]
fn records_high_water_mark_across_reservations() {
let (recording, pool) = pool(1024);
let a = MemoryConsumer::new("a").register(&pool);
let b = MemoryConsumer::new("b").register(&pool);
a.try_grow(300).unwrap();
b.try_grow(400).unwrap();
assert_eq!(recording.peak_reserved(), 700);
a.shrink(300);
b.try_grow(100).unwrap();
assert_eq!(pool.reserved(), 500);
assert_eq!(recording.peak_reserved(), 700);
}
#[test]
fn failed_growth_does_not_move_the_peak() {
let (recording, pool) = pool(1024);
let reservation = MemoryConsumer::new("a").register(&pool);
reservation.try_grow(600).unwrap();
reservation
.try_grow(600)
.expect_err("should exceed the 1024 byte pool");
assert_eq!(recording.peak_reserved(), 600);
}
#[test]
fn reset_clears_the_window_but_not_the_run_maximum() {
let (recording, pool) = pool(1024);
let reservation = MemoryConsumer::new("a").register(&pool);
reservation.try_grow(800).unwrap();
reservation.shrink(800);
recording.reset_peak();
assert_eq!(recording.peak_reserved(), 0);
assert_eq!(recording.max_reserved(), 800);
reservation.try_grow(100).unwrap();
assert_eq!(recording.peak_reserved(), 100);
assert_eq!(recording.max_reserved(), 800);
}
#[test]
fn reset_keeps_what_is_still_reserved() {
let (recording, pool) = pool(1024);
let held = MemoryConsumer::new("held").register(&pool);
held.try_grow(300).unwrap();
recording.reset_peak();
assert_eq!(recording.peak_reserved(), 300);
let query = MemoryConsumer::new("query").register(&pool);
query.try_grow(200).unwrap();
assert_eq!(recording.peak_reserved(), 500);
}
#[test]
fn marks_are_per_instance() {
let (one, one_pool) = pool(1024);
let (two, _two_pool) = pool(1024);
MemoryConsumer::new("a")
.register(&one_pool)
.try_grow(512)
.unwrap();
assert_eq!(one.peak_reserved(), 512);
assert_eq!(two.peak_reserved(), 0);
}
#[test]
fn is_recoverable_from_the_pool_it_is_installed_as() {
let (recording, pool) = pool(1024);
MemoryConsumer::new("a")
.register(&pool)
.try_grow(512)
.unwrap();
let found = PeakRecordingPool::from_pool(&*pool).expect("recorder installed");
assert_eq!(found.peak_reserved(), recording.peak_reserved());
let plain: Arc<dyn MemoryPool> = Arc::new(GreedyMemoryPool::new(1024));
assert!(PeakRecordingPool::from_pool(&*plain).is_none());
}
#[test]
fn delegates_limit_and_name_to_the_wrapped_pool() {
let inner: Arc<dyn MemoryPool> = Arc::new(GreedyMemoryPool::new(4096));
let wrapped = PeakRecordingPool::new(Arc::clone(&inner));
assert_eq!(wrapped.name(), inner.name());
assert_eq!(wrapped.to_string(), inner.to_string());
assert!(matches!(wrapped.memory_limit(), MemoryLimit::Finite(4096)));
}
#[cfg(feature = "arrow_buffer_pool")]
#[test]
fn records_reservations_arriving_through_the_arrow_adapter() {
use crate::memory_pool::arrow::ArrowMemoryPool;
use arrow_buffer::MemoryPool as ArrowMemoryPoolTrait;
let (recording, pool) = pool(4096);
let arrow_pool =
ArrowMemoryPool::new(Arc::clone(&pool), MemoryConsumer::new("arrow"));
let reservation = arrow_pool.reserve(1024);
assert_eq!(pool.reserved(), 1024);
assert_eq!(recording.peak_reserved(), 1024);
drop(reservation);
assert_eq!(pool.reserved(), 0);
assert_eq!(recording.peak_reserved(), 1024);
}
}