use std::{
cell::RefCell,
fmt::Debug,
future::poll_fn,
rc::Rc,
task::{Poll, Waker},
};
use crate::model::{Object, ObjectID, Transition};
#[derive(Clone, Copy)]
enum Op<T> {
Store(T),
Load,
CompareExchange { current: T, new: T },
}
impl<T> Op<T> {
fn is_load(&self) -> bool {
matches!(self, Op::Load)
}
}
struct Request<T> {
transition: Transition,
waker: Waker,
op: Op<T>,
}
struct Record<T> {
transition: Transition,
op: Op<T>,
prev: T,
}
struct Atomic<T> {
value: T,
id: ObjectID,
requests: Vec<Request<T>>,
history: Vec<Record<T>>,
committed: Vec<Option<T>>,
seq: usize,
}
impl<T: Copy + PartialEq + Debug> Atomic<T> {
fn new(id: ObjectID, value: T) -> Self {
Self {
value,
id,
requests: Vec::new(),
history: Vec::new(),
committed: Vec::new(),
seq: 0,
}
}
fn register(&mut self, op: Op<T>, waker: Waker) -> Transition {
let transition = Transition::new(self.id, self.seq);
self.seq += 1;
self.committed.push(None);
self.requests.push(Request {
transition,
waker,
op,
});
transition
}
fn result(&self, seq: usize) -> Option<T> {
self.committed[seq]
}
fn apply(&mut self, t: Transition) {
let Some(i) = self.requests.iter().position(|r| r.transition == t) else {
panic!("transition must be enabled");
};
let req = self.requests.remove(i);
let prev = self.value;
match req.op {
Op::Store(new) => self.value = new,
Op::CompareExchange { current, new } if prev == current => self.value = new,
Op::Load | Op::CompareExchange { .. } => {}
}
self.history.push(Record {
transition: t,
op: req.op,
prev,
});
self.committed[t.seq] = Some(prev);
req.waker.wake();
}
fn enabled(&self) -> Vec<Transition> {
self.requests.iter().map(|r| r.transition).collect()
}
fn enabled_into(&self, out: &mut Vec<Transition>) {
out.extend(self.requests.iter().map(|r| r.transition));
}
fn op_of(&self, t: Transition) -> Op<T> {
let pending = self.requests.iter().map(|r| (r.transition, r.op));
let committed = self.history.iter().map(|r| (r.transition, r.op));
pending
.chain(committed)
.find(|(transition, _)| *transition == t)
.map(|(_, op)| op)
.expect("transition not registered on this atomic")
}
fn depends(&self, t1: Transition, t2: Transition) -> bool {
!(self.op_of(t1).is_load() && self.op_of(t2).is_load())
}
fn label(&self, t: Transition) -> String {
let Some(rec) = self.history.iter().find(|r| r.transition == t) else {
panic!("label called on an unapplied transition");
};
match rec.op {
Op::Store(new) => format!("store {new:?} (was {:?})", rec.prev),
Op::Load => format!("load -> {:?}", rec.prev),
Op::CompareExchange { current, new } if rec.prev == current => {
format!("cas({current:?}->{new:?}) ok")
}
Op::CompareExchange { current, new } => {
format!("cas({current:?}->{new:?}) fail: {:?}", rec.prev)
}
}
}
}
#[derive(Clone)]
pub struct Handle<T: Copy + PartialEq + Debug> {
atomic: Rc<RefCell<Atomic<T>>>,
}
impl<T: Copy + PartialEq + Debug> Handle<T> {
pub(crate) fn new(id: ObjectID, value: T) -> Self {
Self {
atomic: Rc::new(RefCell::new(Atomic::new(id, value))),
}
}
pub async fn store(&self, value: T) -> T {
self.request(Op::Store(value)).await
}
pub async fn load(&self) -> T {
self.request(Op::Load).await
}
pub async fn compare_exchange(&self, current: T, new: T) -> Result<T, T> {
let prev = self.request(Op::CompareExchange { current, new }).await;
if prev == current { Ok(prev) } else { Err(prev) }
}
async fn request(&self, op: Op<T>) -> T {
let mut registered: Option<Transition> = None;
let mut op = Some(op);
poll_fn(move |cx| {
if let Some(t) = registered {
return match self.atomic.borrow().result(t.seq) {
Some(value) => Poll::Ready(value),
None => Poll::Pending,
};
}
let op = op.take().expect("request future polled after completion");
registered = Some(self.atomic.borrow_mut().register(op, cx.waker().clone()));
Poll::Pending
})
.await
}
}
impl<T: Copy + PartialEq + Debug + 'static> Object for Handle<T> {
fn apply(&mut self, t: Transition) {
self.atomic.borrow_mut().apply(t);
}
fn enabled(&self) -> Vec<Transition> {
self.atomic.borrow().enabled()
}
fn enabled_into(&self, out: &mut Vec<Transition>) {
self.atomic.borrow().enabled_into(out);
}
fn may_block(&self) -> bool {
false
}
fn label(&self, t: Transition) -> String {
self.atomic.borrow().label(t)
}
fn depends(&self, t1: Transition, t2: Transition) -> bool {
self.atomic.borrow().depends(t1, t2)
}
}
impl<T: Copy + PartialEq + Debug> std::fmt::Debug for Handle<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self.atomic.try_borrow() {
Ok(a) => f
.debug_struct("Atomic")
.field("id", &a.id)
.finish_non_exhaustive(),
Err(_) => f.debug_struct("Atomic").finish_non_exhaustive(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{Executor, ProcessResult};
use std::cell::Cell;
use std::future::Future;
fn drive(exec: &mut Executor, obj: &mut impl Object) {
exec.execute().unwrap();
while let Some(&t) = obj.enabled().first() {
obj.apply(t);
exec.execute().unwrap();
}
}
fn run_single(atomic: &Handle<u32>, body: impl Future<Output = ProcessResult> + 'static) {
let mut exec = Executor::default();
exec.schedule(body);
drive(&mut exec, &mut atomic.clone());
}
fn slot<T: Copy>(init: T) -> Rc<Cell<T>> {
Rc::new(Cell::new(init))
}
fn enabled_of(atomic: &Handle<u32>, pid: usize) -> Transition {
*atomic.enabled().iter().find(|t| t.pid == pid).unwrap()
}
#[test]
fn load_observes_initial_value() {
let a = Handle::new(0, 42u32);
let seen = slot(0);
let (h, dst) = (a.clone(), seen.clone());
run_single(&a, async move {
dst.set(h.load().await);
Ok(())
});
assert_eq!(seen.get(), 42);
}
#[test]
fn store_writes_and_returns_previous() {
let a = Handle::new(0, 1u32);
let (prev, after) = (slot(0), slot(0));
let (h, p, af) = (a.clone(), prev.clone(), after.clone());
run_single(&a, async move {
p.set(h.store(9).await);
af.set(h.load().await);
Ok(())
});
assert_eq!(prev.get(), 1);
assert_eq!(after.get(), 9);
}
#[test]
fn compare_exchange_swaps_on_match() {
let a = Handle::new(0, 1u32);
let (res, after) = (slot(Err(0)), slot(0));
let (h, r, af) = (a.clone(), res.clone(), after.clone());
run_single(&a, async move {
r.set(h.compare_exchange(1, 9).await);
af.set(h.load().await);
Ok(())
});
assert_eq!(res.get(), Ok(1));
assert_eq!(after.get(), 9);
}
#[test]
fn compare_exchange_observes_commit_time_value() {
let a = Handle::new(0, 1u32);
let res = slot(Ok(0));
let mut exec = Executor::default();
let (cas_h, r) = (a.clone(), res.clone());
exec.schedule(async move {
r.set(cas_h.compare_exchange(1, 9).await);
Ok(())
});
let store_h = a.clone();
exec.schedule(async move {
store_h.store(5).await;
Ok(())
});
exec.execute().unwrap();
let mut obj = a.clone();
obj.apply(enabled_of(&obj, 1));
exec.execute().unwrap();
obj.apply(enabled_of(&obj, 0));
exec.execute().unwrap();
assert_eq!(res.get(), Err(5));
}
}