use std::sync::atomic::Ordering;
use std::ptr;
use std::sync::atomic::AtomicPtr;
pub struct TakeLock<T:Sync>{
inner: AtomicPtr<T>
}
impl<T:Sync> Drop for TakeLock<T>{
fn drop(&mut self) {
let p: *mut T = *self.inner.get_mut();
if p.is_null() {
return;
}
unsafe{
_ = Box::from_raw(p);
}
}
}
impl<T:Sync> Default for TakeLock<T>{
fn default() -> Self {
Self{inner:AtomicPtr::new(ptr::null_mut())}
}
}
fn safe_to_raw<T>(opt: Option<Box<T>>) -> *mut T {
match opt {
Some(b) => Box::into_raw(b),
None => std::ptr::null_mut(),
}
}
unsafe fn raw_to_safe<T>(raw: *mut T) -> Option<Box<T>> {
if raw.is_null(){
None
}else {
unsafe{Some(Box::from_raw(raw))}
}
}
impl<T:Sync> TakeLock<T>{
pub fn new(op:Option<Box<T>>) -> Self {
Self{inner:AtomicPtr::new(safe_to_raw(op))}
}
pub fn is_empty(&self) -> bool{
self.inner.load(Ordering::Acquire).is_null()
}
pub fn take(&self) -> Option<Box<T>> {
self.swap(None)
}
pub fn put(&self,next:Option<Box<T>>) -> Result<(),Option<Box<T>>>{
let p = safe_to_raw(next);
match self.inner.compare_exchange(
ptr::null_mut(),
p,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => Ok(()),
Err(_) => unsafe{Err(raw_to_safe(p))},
}
}
pub fn swap(&self,next:Option<Box<T>>) -> Option<Box<T>>{
let p = safe_to_raw(next);
let mut cur = self.inner.load(Ordering::Acquire);
loop{
match self.inner.compare_exchange(
cur,
p,
Ordering::AcqRel,
Ordering::Acquire
) {
Ok(b) => return unsafe{raw_to_safe(b)},
Err(c) => cur=c,
}
}
}
pub fn into_inner(self) -> Option<Box<T>>{
let p: *mut T = self.inner.load(Ordering::Relaxed);
if p.is_null() {
return None;
}
unsafe{
self.inner.store(ptr::null_mut(),Ordering::Relaxed);
return Some(Box::from_raw(p));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{
atomic::{AtomicUsize, Ordering},
Arc,
};
use std::thread;
#[derive(Debug)]
struct DropCounter {
hits: Arc<AtomicUsize>,
}
impl DropCounter {
fn new(h: Arc<AtomicUsize>) -> Self {
Self { hits: h }
}
}
impl Drop for DropCounter {
fn drop(&mut self) {
self.hits.fetch_add(1, Ordering::Relaxed);
}
}
#[test]
fn put_take_roundtrip() {
let drops = Arc::new(AtomicUsize::new(0));
{
let lock = TakeLock::default();
assert!(lock.is_empty());
lock.put(Some(Box::new(DropCounter::new(drops.clone()))))
.unwrap();
assert!(!lock.is_empty());
let _v = lock.take().unwrap(); assert!(lock.is_empty());
}
assert_eq!(drops.load(Ordering::Relaxed), 1);
}
#[test]
fn put_when_full_fails() {
let lock = TakeLock::new(Some(Box::new(5u32)));
let err = lock.put(Some(Box::new(99u32))).err().unwrap();
assert_eq!(*err.unwrap(), 99); assert_eq!(*lock.take().unwrap(), 5);
}
#[test]
fn swap_returns_previous() {
let lock = TakeLock::new(Some(Box::new(1u8)));
let old = lock.swap(Some(Box::new(2u8))).unwrap();
assert_eq!(*old, 1);
assert_eq!(*lock.take().unwrap(), 2);
}
#[test]
fn only_one_thread_can_put() {
let lock = Arc::new(TakeLock::default());
let l1 = lock.clone();
let l2 = lock.clone();
let t1 = thread::spawn(move || l1.put(Some(Box::new(10i32))).is_ok());
let t2 = thread::spawn(move || l2.put(Some(Box::new(20i32))).is_ok());
let r1 = t1.join().unwrap();
let r2 = t2.join().unwrap();
assert!(r1 ^ r2);
let v = lock.take().unwrap();
assert!(*v == 10 || *v == 20);
}
#[test]
fn into_inner_moves_value_out() {
let lock = TakeLock::new(Some(Box::new(42usize)));
let inner = TakeLock::into_inner(lock).unwrap();
assert_eq!(*inner, 42);
}
#[test]
fn stress_put_take_swap() {
const THREADS: usize = 8;
const ITERS: usize = 250_000;
let drops = Arc::new(AtomicUsize::new(0));
let lock = Arc::new(TakeLock::<DropCounter>::default());
let mut handles = Vec::with_capacity(THREADS);
for tid in 0..THREADS {
let lock = Arc::clone(&lock);
let drops = Arc::clone(&drops);
handles.push(thread::spawn(move || {
for i in 0..ITERS {
match i.wrapping_add(tid) % 10 {
0..=2 => {
let _ = lock.put(Some(Box::new(DropCounter::new(drops.clone()))));
}
3..=5 => {
drop(lock.take());
}
_ => {
let new = if i & 1 == 0 {
Some(Box::new(DropCounter::new(drops.clone())))
} else {
None
};
drop(lock.swap(new));
}
}
}
}));
}
for h in handles { h.join().unwrap(); }
drop(lock.take());
assert!(lock.is_empty(), "lock left non‑empty after stress run");
assert!(
drops.load(Ordering::Relaxed) > 0,
"stress did not move any value through the lock",
);
}
#[test]
fn swap_ping_pong() {
const THREADS: usize = 4;
const ROTATIONS: usize = 1_000_000;
let drops = Arc::new(AtomicUsize::new(0));
let lock = Arc::new(TakeLock::new(
Some(Box::new(DropCounter::new(drops.clone())))
));
let mut handles = Vec::with_capacity(THREADS);
for _ in 0..THREADS {
let lock = Arc::clone(&lock);
let drops = Arc::clone(&drops);
handles.push(thread::spawn(move || {
for _ in 0..ROTATIONS {
let _ = lock.swap(Some(Box::new(DropCounter::new(drops.clone()))));
let _ = lock.swap(None);
}
}));
}
for h in handles { h.join().unwrap(); }
drop(lock.take());
assert!(lock.is_empty());
assert!(drops.load(Ordering::Relaxed) > 0);
}
}