use crate::env_var::config;
use core::marker::PhantomData;
use indexmap::IndexSet;
use parking_lot::{Condvar, Mutex};
use std::collections::BTreeMap;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use tracing::trace;
pub(crate) trait LamellarAlloc {
fn new(id: String) -> Self;
fn init(&mut self, start_addr: usize, size: usize); #[allow(dead_code)]
fn malloc(&self, size: usize, align: usize) -> usize;
fn try_malloc(&self, size: usize, align: usize) -> Option<usize>;
fn fake_malloc(&self, size: usize, align: usize) -> bool;
fn free(&self, addr: usize) -> Result<(), usize>;
fn find(&self, addr: usize) -> Option<usize>;
fn space_avail(&self) -> usize;
fn occupied(&self) -> usize;
}
fn calc_padding(addr: usize, align: usize) -> usize {
let rem = addr % align;
if rem == 0 {
0
} else {
align - rem
}
}
#[derive(Clone, Debug)]
struct FreeEntries {
sizes: BTreeMap<usize, IndexSet<usize>>, addrs: BTreeMap<usize, (usize, usize)>, }
impl FreeEntries {
fn new() -> FreeEntries {
FreeEntries {
sizes: BTreeMap::new(),
addrs: BTreeMap::new(),
}
}
fn merge(&mut self) {
let mut i = 0;
while i < self.addrs.len() - 1 {
let (faddr, (fsize, fpadding)) = self.addrs.pop_first().unwrap();
let (naddr, (nsize, npadding)) = self.addrs.pop_first().unwrap();
if faddr + fsize + fpadding == naddr {
let new_size = fsize + nsize;
let new_padding = fpadding + npadding;
assert!(new_padding == 0);
let new_addr = faddr;
self.remove_size(naddr, nsize);
self.remove_size(faddr, fsize);
self.addrs.insert(new_addr, (new_size, new_padding));
self.sizes
.entry(new_size)
.or_insert(IndexSet::new())
.insert(new_addr);
} else {
self.addrs.insert(faddr, (fsize, fpadding));
self.addrs.insert(naddr, (nsize, npadding));
i += 1;
}
}
}
fn remove_size(&mut self, addr: usize, size: usize) {
let mut remove_size = false;
if let Some(addrs) = self.sizes.get_mut(&size) {
addrs.swap_remove(&addr);
if addrs.is_empty() {
remove_size = true;
}
}
if remove_size {
self.sizes.remove(&size);
}
}
}
#[derive(Clone, Debug)]
pub(crate) struct BTreeAlloc {
free_entries: Arc<(Mutex<FreeEntries>, Condvar)>,
allocated_addrs: Arc<(Mutex<BTreeMap<usize, (usize, usize)>>, Condvar)>, pub(crate) start_addr: usize,
pub(crate) max_size: usize,
id: String,
free_space: Arc<AtomicUsize>,
}
impl BTreeAlloc {}
impl LamellarAlloc for BTreeAlloc {
fn new(id: String) -> BTreeAlloc {
BTreeAlloc {
free_entries: Arc::new((Mutex::new(FreeEntries::new()), Condvar::new())),
allocated_addrs: Arc::new((Mutex::new(BTreeMap::new()), Condvar::new())),
start_addr: 0,
max_size: 0,
id,
free_space: Arc::new(AtomicUsize::new(0)),
}
}
fn init(&mut self, start_addr: usize, size: usize) {
self.start_addr = start_addr;
self.max_size = size;
let &(ref lock, ref _cvar) = &*self.free_entries;
let mut free_entries = lock.lock();
let mut temp = IndexSet::new();
temp.insert(start_addr);
free_entries.sizes.insert(size, temp);
free_entries.addrs.insert(start_addr, (size, 0));
self.free_space.store(size, Ordering::SeqCst);
}
fn malloc(&self, size: usize, align: usize) -> usize {
let mut val = self.try_malloc(size, align);
let mut timer = std::time::Instant::now();
while let None = val {
val = self.try_malloc(size, align);
if timer.elapsed().as_secs_f64() > config().deadlock_warning_timeout {
println!("[WARNING] Potential deadlock detected when trying to allocate more memory.\n\
The deadlock timeout can be set via the LAMELLAR_DEADLOCK_WARNING_TIMEOUT environment variable, the current timeout is {} seconds\n\
To view backtrace set RUST_LIB_BACKTRACE=1\n\
{}",config().deadlock_warning_timeout,std::backtrace::Backtrace::capture());
timer = std::time::Instant::now();
}
}
val.unwrap()
}
fn try_malloc(&self, size: usize, align: usize) -> Option<usize> {
let &(ref lock, ref cvar) = &*self.free_entries;
let mut free_entries = lock.lock();
let mut addr: Option<usize> = None;
let mut remove_size: Option<usize> = None;
let upper_size = size + align - 1;
let mut try_again = true;
while try_again {
if let Some((free_size, addrs)) = free_entries.sizes.range_mut(upper_size..).next() {
addr = addrs.pop();
if addrs.is_empty() {
remove_size = Some(free_size.clone());
}
if let Some(a) = addr {
let padding = calc_padding(a, align);
let full_size = size + padding;
if let Some((fsize, fpadding)) = free_entries.addrs.remove(&a) {
if fsize + fpadding != full_size {
let remaining = (fsize + fpadding) - full_size;
let new_addr = a + full_size;
free_entries
.sizes
.entry(remaining)
.or_insert(IndexSet::new()) .insert(new_addr); free_entries.addrs.insert(new_addr, (remaining, 0));
}
} else {
panic!("{:?} addr {:?} not found in free_entries", self.id, a);
}
}
try_again = false;
} else {
cvar.wait_for(&mut free_entries, std::time::Duration::from_millis(1));
free_entries.merge();
if free_entries.sizes.range_mut(upper_size..).next().is_none() {
try_again = false;
}
}
}
if let Some(rsize) = remove_size {
free_entries.sizes.remove(&rsize);
}
drop(free_entries);
addr = if let Some(a) = addr {
let padding = calc_padding(a, align);
let full_size = size + padding;
let &(ref lock, ref _cvar) = &*self.allocated_addrs;
let mut allocated_addrs = lock.lock();
allocated_addrs.insert(a + padding, (size, padding));
trace!(target: "ucx",
"alloc addr 0x{:x} = 0x{a:x} + padding 0x{padding:x} ({padding}) (size: {size}, align: {align}) {:?}",
a + padding,
self.free_space.load(Ordering::SeqCst)
);
self.free_space.fetch_sub(full_size, Ordering::SeqCst);
let new_addr = a + padding;
Some(new_addr)
} else {
None
};
addr
}
fn fake_malloc(&self, size: usize, align: usize) -> bool {
let &(ref lock, ref _cvar) = &*self.free_entries;
let mut free_entries = lock.lock();
let upper_size = size + align - 1; if let Some((_, _)) = free_entries.sizes.range_mut(upper_size..).next() {
return true;
} else {
free_entries.merge();
if let Some((_, _)) = free_entries.sizes.range_mut(upper_size..).next() {
return true;
}
return false;
}
}
fn free(&self, addr: usize) -> Result<(), usize> {
let &(ref lock, ref _cvar) = &*self.allocated_addrs;
let mut allocated_addrs = lock.lock();
if let Some((size, padding)) = allocated_addrs.remove(&addr) {
let full_size = size + padding;
self.free_space.fetch_add(full_size, Ordering::SeqCst);
drop(allocated_addrs);
let unpadded_addr = addr - padding;
let full_size = size + padding;
let mut temp_addr = unpadded_addr;
let mut temp_size = full_size;
let mut remove = Vec::new();
let &(ref lock, ref cvar) = &*self.free_entries;
let mut free_entries = lock.lock();
if let Some((faddr, (fsize, fpadding))) =
free_entries.addrs.range(..temp_addr).next_back()
{
if faddr + fsize + fpadding == addr {
temp_addr = *faddr;
temp_size += fsize + fpadding;
remove.push((*faddr, *fsize, *fpadding));
}
}
if let Some((faddr, (fsize, fpadding))) = free_entries.addrs.range(addr..).next() {
if temp_addr + temp_size == *faddr {
temp_size += fsize + fpadding;
remove.push((*faddr, *fsize, *fpadding));
}
}
for (raddr, rsize, rpadding) in remove {
let rfull_size = rsize + rpadding;
free_entries.addrs.remove(&raddr);
free_entries.remove_size(raddr, rfull_size);
}
free_entries.addrs.insert(temp_addr, (temp_size, 0));
free_entries
.sizes
.entry(temp_size)
.or_insert(IndexSet::new())
.insert(temp_addr);
cvar.notify_all();
Ok(())
} else {
Err(addr)
}
}
fn find(&self, addr: usize) -> Option<usize> {
let &(ref lock, ref _cvar) = &*self.allocated_addrs;
let allocated_addrs = lock.lock();
if let Some((size, _padding)) = allocated_addrs.get(&addr) {
trace!(target: "ucx",
"find addr 0x:{:x} size: {size} padding: {_padding} free_space: {}",
addr,
self.free_space.load(Ordering::SeqCst)
);
return Some(*size);
}
None
}
fn space_avail(&self) -> usize {
self.free_space.load(Ordering::SeqCst)
}
fn occupied(&self) -> usize {
self.max_size - self.free_space.load(Ordering::SeqCst)
}
}
#[derive(Clone)]
#[allow(dead_code)]
pub(crate) struct ObjAlloc<T: Copy> {
free_entries: Arc<(Mutex<Vec<usize>>, Condvar)>,
start_addr: usize,
max_size: usize,
num_entries: usize,
_id: String,
phantom: PhantomData<T>,
}
impl<T: Copy> LamellarAlloc for ObjAlloc<T> {
fn new(id: String) -> ObjAlloc<T> {
ObjAlloc {
free_entries: Arc::new((Mutex::new(Vec::new()), Condvar::new())),
start_addr: 0,
max_size: 0,
num_entries: 0,
_id: id,
phantom: PhantomData,
}
}
fn init(&mut self, start_addr: usize, size: usize) {
let align = std::mem::align_of::<T>();
let padding = calc_padding(start_addr, align);
self.start_addr = start_addr + padding;
self.max_size = size;
let &(ref lock, ref _cvar) = &*self.free_entries;
let mut free_entries = lock.lock();
*free_entries = ((start_addr + padding)..(start_addr + size))
.step_by(std::mem::size_of::<T>())
.collect();
self.num_entries = free_entries.len();
}
fn malloc(&self, size: usize, align: usize) -> usize {
let mut val = self.try_malloc(size, align);
while let None = val {
val = self.try_malloc(size, align);
}
val.unwrap()
}
fn try_malloc(&self, size: usize, _align: usize) -> Option<usize> {
assert_eq!(
size, 1,
"ObjAlloc does not currently support multiobject allocations"
);
let &(ref lock, ref cvar) = &*self.free_entries;
let mut free_entries = lock.lock();
if let Some(addr) = free_entries.pop() {
return Some(addr);
} else {
cvar.wait_for(&mut free_entries, std::time::Duration::from_millis(1));
return None;
}
}
fn fake_malloc(&self, size: usize, _align: usize) -> bool {
assert_eq!(
size, 1,
"ObjAlloc does not currently support multiobject allocations"
);
let &(ref lock, ref _cvar) = &*self.free_entries;
let free_entries = lock.lock();
if free_entries.len() > 1 {
true
} else {
false
}
}
fn free(&self, addr: usize) -> Result<(), usize> {
let &(ref lock, ref cvar) = &*self.free_entries;
let mut free_entries = lock.lock();
free_entries.push(addr);
cvar.notify_all();
Ok(())
}
fn find(&self, addr: usize) -> Option<usize> {
assert!(addr < self.num_entries);
let &(ref lock, ref _cvar) = &*self.free_entries;
let free_entries = lock.lock();
match free_entries.iter().position(|x| *x == addr) {
Some(_i) => None, None => Some(1),
}
}
fn space_avail(&self) -> usize {
let &(ref lock, ref _cvar) = &*self.free_entries;
let free_entries = lock.lock();
free_entries.len()
}
fn occupied(&self) -> usize {
self.max_size - self.space_avail()
}
}
#[cfg(test)]
mod tests {
use super::*;
use rand::seq::SliceRandom;
use rand::{rngs::StdRng, Rng, SeedableRng};
fn test_malloc(alloc: &mut impl LamellarAlloc) {
alloc.init(0, 100);
let mut rng = StdRng::seed_from_u64(0 as u64);
let mut shuffled: Vec<usize> = (1..11).collect();
shuffled.shuffle(&mut rng);
let mut cnt = 0;
for i in shuffled {
assert_eq!(alloc.malloc(i, 1), cnt);
cnt += i;
}
for i in (cnt..100).step_by(5) {
assert_eq!(alloc.malloc(5, 1), i);
}
}
fn stress<T: LamellarAlloc + Clone + Send + 'static>(alloc: T) {
let mut threads = Vec::new();
let start = std::time::Instant::now();
for _i in 0..10 {
let alloc_clone = alloc.clone();
let t = std::thread::spawn(move || {
let mut rng = rand::rng();
let mut addrs: Vec<usize> = Vec::new();
let mut i = 0;
while i < 100000 {
if rng.random_range(0..2) == 0 || addrs.is_empty() {
if let Some(addr) = alloc_clone.try_malloc(1, 1) {
addrs.push(addr);
i += 1;
}
} else {
let index = rng.random_range(0..addrs.len());
let addr = addrs.remove(index);
alloc_clone
.free(addr)
.expect("Address should have been found and freed");
}
}
for addr in addrs {
alloc_clone
.free(addr)
.expect("Address should have been found and freed");
}
});
threads.push(t);
}
for t in threads {
t.join().unwrap();
}
let time = start.elapsed().as_secs_f64();
println!("time: {:?}", time);
}
#[test]
fn test_bttreealloc_malloc() {
let mut alloc = BTreeAlloc::new("bttree_malloc".to_string());
test_malloc(&mut alloc);
}
#[test]
fn test_bttreealloc_stress() {
let mut alloc = BTreeAlloc::new("bttree_stress".to_string());
alloc.init(0, 100000);
stress(alloc.clone());
let &(ref lock, ref _cvar) = &*alloc.free_entries;
let free_entries = lock.lock();
assert_eq!(free_entries.sizes.len(), 1);
assert_eq!(free_entries.addrs.len(), 1);
let &(ref lock, ref _cvar) = &*alloc.allocated_addrs;
let allocated_addrs = lock.lock();
assert_eq!(allocated_addrs.len(), 0);
}
#[test]
fn test_obj_malloc() {
let mut alloc = ObjAlloc::<u8>::new("obj_malloc_u8".to_string());
alloc.init(0, 10);
for i in 0..10 {
assert_eq!(alloc.malloc(1, 1), 9 - i); }
let mut alloc = ObjAlloc::<u16>::new("obj_malloc_u16".to_string());
alloc.init(0, 10); for i in 0..5 {
assert_eq!(alloc.malloc(1, 1), 8 - (i * std::mem::size_of::<u16>()));
}
assert_eq!(alloc.try_malloc(1, 1), None);
}
#[test]
fn test_obj_u8_stress() {
let mut alloc = ObjAlloc::<u8>::new("obj_malloc_u8".to_string());
alloc.init(0, 100000);
stress(alloc.clone());
let &(ref lock, ref _cvar) = &*alloc.free_entries;
let free_entries = lock.lock();
assert_eq!(free_entries.len(), 100000 / std::mem::size_of::<u8>());
}
#[test]
fn test_obj_f64_stress() {
let mut alloc = ObjAlloc::<f64>::new("obj_malloc_u8".to_string());
alloc.init(0, 100000);
stress(alloc.clone());
let &(ref lock, ref _cvar) = &*alloc.free_entries;
let free_entries = lock.lock();
assert_eq!(free_entries.len(), 100000 / std::mem::size_of::<f64>());
}
}