#![allow(non_camel_case_types)]
use std::sync::{Arc, RwLock, RwLockReadGuard, RwLockWriteGuard};
pub struct iBag<T: Sized> {
inner: Arc<RwLock<T>>,
}
impl<T> iBag<T> where T: Sized {
pub fn new(value: T) -> Self {
Self {
inner: Arc::new(RwLock::new(value)),
}
}
pub fn load(&self) -> RwLockReadGuard<T> {
self.inner.read().unwrap()
}
pub fn write(&self) -> RwLockWriteGuard<T> {
self.inner.write().unwrap()
}
pub fn with<F, R>(&self, f: F) -> R
where
F: FnOnce(&mut T) -> R,
{
let mut guard = self.write();
f(&mut *guard)
}
pub fn with_read<F, R>(&self, f: F) -> R
where
F: FnOnce(&T) -> R,
{
let guard = self.load();
f(&*guard)
}
}
impl<T: Sized> Clone for iBag<T> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
unsafe impl<T> Send for iBag<T> {}
unsafe impl<T> Sync for iBag<T> {}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::{thread, vec};
#[test]
fn test_basic_operations() {
let bag = iBag::new(42);
assert_eq!(unsafe { *bag.load() }, 42);
bag.with(|val| {
*val = 100;
});
assert_eq!(unsafe { *bag.load() }, 100);
}
#[test]
fn test_with_closure() {
let bag = iBag::new(String::from("test"));
let len = bag.with(|s| s.len());
assert_eq!(len, 4);
}
#[test]
fn test_clone() {
let bag1 = iBag::new(42);
let bag2 = bag1.clone();
unsafe {
let b1 = *bag1.load();
let b2 = *bag2.load();
assert_eq!(b1, b2);
}
}
#[test]
fn test_thread_safety() {
let bag = Arc::new(iBag::new(0));
let mut handles = vec![];
for _ in 0..10 {
let bag = bag.clone();
handles.push(thread::spawn(move || {
for _ in 0..1000 {
let val = unsafe { *bag.load() };
bag.with(|v| {
*v = val + 1;
});
}
}));
}
for handle in handles {
handle.join().unwrap();
}
}
#[test]
fn test_send_sync() {
let bag: iBag<usize> = iBag::new(42);
let _ = Arc::new(bag);
}
#[test]
fn test_drop() {
let bag = iBag::new(42);
drop(bag);
}
#[test]
fn test_thread_safety_with_drop() {
let bag = iBag::new(0);
(0..4).for_each(|i| {
let b = bag.clone();
thread::spawn(move || {
let r = b.with(|v| {
*v = i;
*v
});
assert_eq!(r, i);
});
});
}
#[test]
fn test_thread_safety_with_struct() {
struct inner {
pub a: i32,
pub b: i32,
pub c: String,
}
let iv = inner {
a: 0,
b: 0,
c: String::from("test"),
};
let bag = iBag::new(iv);
(0..4).for_each(|i| {
let b = bag.clone();
let mut handles = vec![];
handles.push(thread::spawn(move || {
println!("thread: {}", i);
let r = b.with(|v| {
(*v).a = i;
(*v).b = i+1;
});
let r = b.load();
unsafe {
assert_eq!((*r).a, i);
assert_eq!((*r).b, i+1);
}
}));
for handle in handles {
handle.join().unwrap();
}
});
}
}