use mpmc::{LockFreeQueue};
use std::isize;
use std::thread;
use std::sync::atomic::{Ordering, AtomicIsize};
const DISCONNECTED: isize = isize::MIN;
const FUDGE: isize = 1024;
pub struct Canal<T> {
queue: LockFreeQueue<T>,
cnt: AtomicIsize,
channels: AtomicIsize,
ports: AtomicIsize,
sender_drain: AtomicIsize,
}
#[derive(Debug)]
pub enum Failure {
Empty,
Disconnected,
}
impl<T> Canal<T> {
pub fn new(cap: usize) -> Canal<T> {
Canal {
queue: LockFreeQueue::with_capacity(cap),
cnt: AtomicIsize::new(0),
channels: AtomicIsize::new(1),
ports: AtomicIsize::new(1),
sender_drain: AtomicIsize::new(0),
}
}
}
impl<T: Send> Canal<T> {
pub fn send(&self, t: T) -> Result<(), T> {
if self.ports.load(Ordering::SeqCst) == 0 { return Err(t) }
if self.cnt.load(Ordering::SeqCst) < DISCONNECTED + FUDGE {
return Err(t)
}
try!(self.queue.push(t));
match self.cnt.fetch_add(1, Ordering::SeqCst) {
n if n < DISCONNECTED + FUDGE => {
self.cnt.store(DISCONNECTED, Ordering::SeqCst);
if self.sender_drain.fetch_add(1, Ordering::SeqCst) == 0 {
loop {
loop {
match self.queue.pop() {
Some(..) => {}
None => break,
}
}
if self.sender_drain.fetch_sub(1, Ordering::SeqCst) == 1 {
break
}
}
}
}
n => { assert!(n >= 0) }
}
Ok(())
}
fn decrement(&self) {
match self.cnt.fetch_sub(1, Ordering::SeqCst) {
DISCONNECTED => { self.cnt.store(DISCONNECTED, Ordering::SeqCst); }
n => {
assert!(n >= 0);
}
}
}
pub fn recv(&mut self) -> Result<T, Failure> {
loop {
match self.try_recv() {
Err(Failure::Empty) => {}
data => { return data },
}
thread::yield_now();
}
}
pub fn try_recv(&self) -> Result<T, Failure> {
match self.queue.pop() {
Some(data) => {
self.decrement();
Ok(data)
}
None => {
match self.cnt.load(Ordering::SeqCst) {
n if n != DISCONNECTED => Err(Failure::Empty),
_ => {
match self.queue.pop() {
Some(t) => Ok(t), None => Err(Failure::Disconnected),
}
}
}
}
}
}
pub fn clone_chan(&self) {
self.channels.fetch_add(1, Ordering::SeqCst);
}
pub fn clone_port(&self) {
self.ports.fetch_add(1, Ordering::SeqCst);
}
pub fn drop_chan(&self) {
match self.channels.fetch_sub(1, Ordering::SeqCst) {
1 => {}
n if n > 1 => return,
n => panic!("bad number of channels left {}", n),
}
match self.cnt.swap(DISCONNECTED, Ordering::SeqCst) {
DISCONNECTED => {}
n => { assert!(n >= 0); }
}
}
pub fn drop_port(&self) {
match self.ports.fetch_sub(1, Ordering::SeqCst) {
1 => {}
n if n > 1 => return,
n => panic!("bad number of channels left {}", n),
}
let mut steals = 0;
while {
let cnt = self.cnt.compare_and_swap(steals, DISCONNECTED, Ordering::SeqCst);
cnt != DISCONNECTED && cnt != steals
} {
loop {
match self.queue.pop() {
Some(..) => { steals += 1; }
None => break,
}
}
}
}
}
impl<T> Drop for Canal<T> {
fn drop(&mut self) {
assert_eq!(self.cnt.load(Ordering::SeqCst), DISCONNECTED);
assert_eq!(self.channels.load(Ordering::SeqCst), 0);
assert_eq!(self.ports.load(Ordering::SeqCst), 0);
}
}
#[cfg(test)]
mod tests {
use super::Canal;
use mpmc::mpmc_channel;
use std::thread;
use std::sync::{Arc, Barrier};
#[test]
fn test_send_recv() {
let mut canal = Canal::new(20);
for i in 0..20 {
assert!(canal.send(i as u8).is_ok());
}
for _i in 0..20 {
let popped = canal.recv().unwrap();
let mut found = false;
for x in 0..20 {
if popped == x {
found = true
}
}
assert!(found);
}
canal.drop_port();
canal.drop_chan();
}
#[test]
fn test_send_full() {
let canal = Canal::new(2);
assert!(canal.send(1 as u8).is_ok());
assert!(canal.send(2 as u8).is_ok());
assert!(canal.send(3 as u8).is_err());
canal.drop_port();
canal.drop_chan();
}
#[test]
fn test_blocking() {
let (sn, rc) = mpmc_channel(5);
let barrier = Arc::new(Barrier::new(6));
let mut recv_vec = Vec::new();
for _i in 0..5 {
let b = barrier.clone();
let r = rc.clone();
recv_vec.push(thread::spawn(move || {
b.wait();
let popped = r.recv().expect("Could not recv");
let mut found = false;
for x in 0..5 {
if popped == x {
found = true
}
}
assert!(found);
}));
}
barrier.wait(); thread::yield_now(); thread::spawn(move || {
for i in 0..5 {
sn.send(i as u8).unwrap();
thread::yield_now();
}
for thr in recv_vec.into_iter() {
thr.join().expect("recv thread errored");
}
}).join().expect("send thread errored");;
}
}