use crate::waiter::Waiter;
use std::cell::Cell;
use std::fmt;
use std::io;
use std::marker::PhantomPinned;
use std::num::NonZeroUsize;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
thread_local! {static TAG: Cell<usize> = const { Cell::new(0) }}
#[derive(Debug)]
pub struct ID(NonZeroUsize);
impl ID {
pub unsafe fn from_usize(id: usize) -> Self {
ID(NonZeroUsize::new(id).expect("id should not be zero"))
}
}
impl From<ID> for usize {
fn from(id: ID) -> Self {
id.0.get()
}
}
#[derive(Debug)]
pub struct Error;
pub struct TokenWaiter<T> {
waiter: Waiter<T>,
key: AtomicUsize,
_phantom: PhantomPinned,
}
impl<T> TokenWaiter<T> {
pub fn new() -> Self {
TokenWaiter {
key: AtomicUsize::new(0),
waiter: Waiter::new(),
_phantom: PhantomPinned,
}
}
pub fn id(&self) -> Result<ID, Error> {
let id = self.key.load(Ordering::Relaxed);
if id != 0 {
return Err(Error);
}
let address = self as *const _ as usize;
let tag = TAG.with(|t| {
let x = t.get();
t.set(x + 1);
(x & 0x1f) << 1
});
let id = (address << 3) | tag;
self.key.store(id, Ordering::Relaxed);
Ok(ID(NonZeroUsize::new(id).unwrap()))
}
fn from_id(id: &ID) -> Option<&Self> {
let id = id.0.get();
let address = (id >> 3) & !0x7;
let waiter = unsafe { &*(address as *const Self) };
if waiter
.key
.compare_exchange(id, id + 1, Ordering::AcqRel, Ordering::Relaxed)
.is_ok()
{
Some(waiter)
} else {
None
}
}
pub fn wait_rsp<D: Into<Option<Duration>>>(&self, timeout: D) -> io::Result<T> {
self.waiter.wait_rsp(timeout)
}
pub fn set_rsp(id: ID, rsp: T) {
if let Some(waiter) = Self::from_id(&id) {
waiter.key.store(0, Ordering::Release);
waiter.waiter.set_rsp(rsp);
}
}
}
impl<T> fmt::Debug for TokenWaiter<T> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "TokenWaiter{{ ... }}")
}
}
impl<T> Default for TokenWaiter<T> {
fn default() -> Self {
TokenWaiter::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use may::go;
#[test]
fn token_waiter_id() {
let waiter = TokenWaiter::<usize>::new();
assert!(waiter.id().is_ok());
assert!(waiter.id().is_err());
}
#[test]
fn token_waiter() {
for j in 0..100 {
let result = go!(move || {
let waiter = TokenWaiter::<usize>::new();
let id = waiter.id().unwrap();
go!(move || TokenWaiter::set_rsp(id, j + 100));
assert_eq!(waiter.wait_rsp(None).unwrap(), j + 100);
let id = waiter.id().unwrap();
go!(move || TokenWaiter::set_rsp(id, j));
waiter.wait_rsp(std::time::Duration::from_secs(2)).unwrap()
})
.join()
.unwrap();
assert_eq!(result, j);
}
}
#[test]
fn token_waiter_timeout() {
let result = go!(|| {
let waiter = TokenWaiter::<usize>::new();
let id = waiter.id().unwrap();
let h = go!(move || {
may::coroutine::sleep(Duration::from_millis(102));
TokenWaiter::set_rsp(id, 42)
});
let ret = waiter.wait_rsp(Duration::from_millis(100));
h.join().unwrap();
ret
})
.join()
.unwrap();
assert!(result.is_err());
}
}