use crate::misc_helpers::Overlaps;
use crate::vector_select::FutureVector;
use std::ops::Range;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use tokio::sync::oneshot;
#[derive(Debug, Default)]
pub struct CommonStorageHelper {
weak_write_blockers: std::sync::RwLock<RangeBlockedList>,
strong_write_blockers: std::sync::RwLock<RangeBlockedList>,
}
#[derive(Debug, Default)]
struct RangeBlockedList {
blocked: Vec<Arc<RangeBlocked>>,
}
#[derive(Debug)]
struct RangeBlocked {
range: Range<u64>,
waitlist: std::sync::Mutex<Vec<oneshot::Sender<()>>>,
index: AtomicUsize,
}
#[derive(Debug)]
pub struct RangeBlockedGuard<'a> {
list: &'a std::sync::RwLock<RangeBlockedList>,
block: Option<Arc<RangeBlocked>>,
}
impl CommonStorageHelper {
pub async fn weak_write_blocker(&self, range: Range<u64>) -> RangeBlockedGuard<'_> {
let mut intersecting = FutureVector::new();
let range_block = {
let mut weak = self.weak_write_blockers.write().unwrap();
let strong = self.strong_write_blockers.read().unwrap();
strong.collect_intersecting_await_futures(&range, &mut intersecting);
weak.block(range)
};
intersecting.discarding_join().await.unwrap();
RangeBlockedGuard {
list: &self.weak_write_blockers,
block: Some(range_block),
}
}
pub async fn strong_write_blocker(&self, range: Range<u64>) -> RangeBlockedGuard<'_> {
let mut intersecting = FutureVector::new();
let range_block = {
let weak = self.weak_write_blockers.read().unwrap();
let mut strong = self.strong_write_blockers.write().unwrap();
weak.collect_intersecting_await_futures(&range, &mut intersecting);
strong.collect_intersecting_await_futures(&range, &mut intersecting);
strong.block(range)
};
intersecting.discarding_join().await.unwrap();
RangeBlockedGuard {
list: &self.strong_write_blockers,
block: Some(range_block),
}
}
}
impl RangeBlockedList {
fn collect_intersecting_await_futures(
&self,
check_range: &Range<u64>,
future_vector: &mut FutureVector<(), oneshot::error::RecvError, oneshot::Receiver<()>>,
) {
for range_block in self.blocked.iter() {
if range_block.range.overlaps(check_range) {
let (s, r) = oneshot::channel::<()>();
range_block.waitlist.lock().unwrap().push(s);
future_vector.push(r);
}
}
}
fn block(&mut self, range: Range<u64>) -> Arc<RangeBlocked> {
let range_block = Arc::new(RangeBlocked {
range,
waitlist: Default::default(),
index: self.blocked.len().into(),
});
self.blocked.push(Arc::clone(&range_block));
range_block
}
}
impl Drop for RangeBlockedGuard<'_> {
fn drop(&mut self) {
let block = self.block.take().unwrap();
{
let mut list = self.list.write().unwrap();
let i = block.index.load(Ordering::Relaxed);
let removed = list.blocked.swap_remove(i);
debug_assert!(Arc::ptr_eq(&removed, &block));
if let Some(block) = list.blocked.get(i) {
block.index.store(i, Ordering::Relaxed);
}
}
let block = Arc::into_inner(block).unwrap();
let waitlist = block.waitlist.into_inner().unwrap();
for waiting in waitlist {
waiting.send(()).unwrap();
}
}
}