use super::Borrow;
use crate::park::Unpark;
use std::cell::UnsafeCell;
use std::fmt::{self, Debug};
use std::future::Future;
use std::mem;
use std::pin::Pin;
use std::ptr;
use std::sync::atomic::Ordering::{AcqRel, Acquire, Relaxed, Release, SeqCst};
use std::sync::atomic::{AtomicBool, AtomicPtr, AtomicUsize};
use std::sync::{Arc, Weak};
use std::task::{Context, Poll, RawWaker, RawWakerVTable, Waker};
use std::thread;
use std::usize;
pub(crate) struct Scheduler<U> {
inner: Arc<Inner<U>>,
nodes: List<U>,
}
struct List<U> {
len: usize,
head: *const Node<U>,
tail: *const Node<U>,
}
struct Inner<U> {
unpark: U,
tick_num: AtomicUsize,
head_readiness: AtomicPtr<Node<U>>,
tail_readiness: UnsafeCell<*const Node<U>>,
stub: Arc<Node<U>>,
}
unsafe impl<U: Sync + Send> Send for Inner<U> {}
unsafe impl<U: Sync + Send> Sync for Inner<U> {}
struct Node<U> {
item: UnsafeCell<Option<Task>>,
notified_at: AtomicUsize,
next_all: UnsafeCell<*const Node<U>>,
prev_all: UnsafeCell<*const Node<U>>,
next_readiness: AtomicPtr<Node<U>>,
queued: AtomicBool,
queue: Weak<Inner<U>>,
}
enum Dequeue<U> {
Data(*const Node<U>),
Empty,
Yield,
Inconsistent,
}
struct Task(Pin<Box<dyn Future<Output = ()>>>);
pub(crate) struct Scheduled<'a, U> {
task: &'a mut Task,
node: &'a Arc<Node<U>>,
done: &'a mut bool,
}
impl<U> Scheduler<U>
where
U: Unpark,
{
pub(crate) fn new(unpark: U) -> Self {
let stub = Arc::new(Node {
item: UnsafeCell::new(None),
notified_at: AtomicUsize::new(0),
next_all: UnsafeCell::new(ptr::null()),
prev_all: UnsafeCell::new(ptr::null()),
next_readiness: AtomicPtr::new(ptr::null_mut()),
queued: AtomicBool::new(true),
queue: Weak::new(),
});
let stub_ptr = &*stub as *const Node<U>;
let inner = Arc::new(Inner {
unpark,
tick_num: AtomicUsize::new(0),
head_readiness: AtomicPtr::new(stub_ptr as *mut _),
tail_readiness: UnsafeCell::new(stub_ptr),
stub,
});
Scheduler {
inner,
nodes: List::new(),
}
}
pub(crate) fn waker(&self) -> Waker {
waker_inner(self.inner.clone())
}
pub(crate) fn schedule(&mut self, item: Pin<Box<dyn Future<Output = ()>>>) {
let tick_num = self.inner.tick_num.load(SeqCst);
let node = Arc::new(Node {
item: UnsafeCell::new(Some(Task::new(item))),
notified_at: AtomicUsize::new(tick_num),
next_all: UnsafeCell::new(ptr::null_mut()),
prev_all: UnsafeCell::new(ptr::null_mut()),
next_readiness: AtomicPtr::new(ptr::null_mut()),
queued: AtomicBool::new(true),
queue: Arc::downgrade(&self.inner),
});
let ptr = self.nodes.push_back(node);
self.inner.enqueue(ptr);
}
pub(crate) fn has_pending_futures(&mut self) -> bool {
unsafe { self.inner.has_pending_futures() }
}
pub(crate) fn tick(&mut self, eid: u64, num_futures: &AtomicUsize) -> bool {
let mut ret = false;
let tick = self.inner.tick_num.fetch_add(1, SeqCst).wrapping_add(1);
loop {
let node = match unsafe { self.inner.dequeue(Some(tick)) } {
Dequeue::Empty => {
return ret;
}
Dequeue::Yield => {
self.inner.unpark.unpark();
return ret;
}
Dequeue::Inconsistent => {
thread::yield_now();
continue;
}
Dequeue::Data(node) => node,
};
ret = true;
debug_assert!(node != self.inner.stub());
unsafe {
if (*(*node).item.get()).is_none() {
let node = ptr2arc(node);
assert!((*node.next_all.get()).is_null());
assert!((*node.prev_all.get()).is_null());
continue;
};
struct Bomb<'a, U: Unpark> {
borrow: &'a mut Borrow<'a, U>,
node: Option<Arc<Node<U>>>,
}
impl<U: Unpark> Drop for Bomb<'_, U> {
fn drop(&mut self) {
if let Some(node) = self.node.take() {
self.borrow.enter(|| release_node(node))
}
}
}
let node = self.nodes.remove(node);
let mut borrow = Borrow {
id: eid,
scheduler: self,
num_futures,
};
let mut bomb = Bomb {
node: Some(node),
borrow: &mut borrow,
};
let mut done = false;
{
let node = bomb.node.as_ref().unwrap();
let item = (*node.item.get()).as_mut().unwrap();
let prev = (*node).queued.swap(false, SeqCst);
assert!(prev);
let borrow = &mut *bomb.borrow;
let mut scheduled = Scheduled {
task: item,
node: bomb.node.as_ref().unwrap(),
done: &mut done,
};
if borrow.enter(|| scheduled.tick()) {
borrow.num_futures.fetch_sub(2, SeqCst);
}
}
if !done {
let node = bomb.node.take().unwrap();
bomb.borrow.scheduler.nodes.push_back(node);
}
}
}
}
}
impl<U: Unpark> Scheduled<'_, U> {
pub(crate) fn tick(&mut self) -> bool {
let waker = unsafe {
waker_ref(self.node)
};
let mut cx = Context::from_waker(&waker);
let ret = match self.task.0.as_mut().poll(&mut cx) {
Poll::Ready(()) => true,
Poll::Pending => false,
};
*self.done = ret;
ret
}
}
impl Task {
pub(crate) fn new(future: Pin<Box<dyn Future<Output = ()> + 'static>>) -> Self {
Task(future)
}
}
impl fmt::Debug for Task {
fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
fmt.debug_struct("Task").finish()
}
}
fn release_node<U>(node: Arc<Node<U>>) {
let prev = node.queued.swap(true, SeqCst);
unsafe {
drop((*node.item.get()).take());
}
if prev {
mem::forget(node);
}
}
impl<U> Debug for Scheduler<U> {
fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(fmt, "Scheduler {{ ... }}")
}
}
impl<U> Drop for Scheduler<U> {
fn drop(&mut self) {
while let Some(node) = self.nodes.pop_front() {
release_node(node);
}
}
}
impl<U> Inner<U> {
fn enqueue(&self, node: *const Node<U>) {
unsafe {
debug_assert!((*node).queued.load(Relaxed));
(*node).next_readiness.store(ptr::null_mut(), Relaxed);
let node = node as *mut _;
let prev = self.head_readiness.swap(node, AcqRel);
(*prev).next_readiness.store(node, Release);
}
}
unsafe fn has_pending_futures(&self) -> bool {
let tail = *self.tail_readiness.get();
let next = (*tail).next_readiness.load(Acquire);
if tail == self.stub() && next.is_null() {
return false;
}
true
}
unsafe fn dequeue(&self, tick: Option<usize>) -> Dequeue<U> {
let mut tail = *self.tail_readiness.get();
let mut next = (*tail).next_readiness.load(Acquire);
if tail == self.stub() {
if next.is_null() {
return Dequeue::Empty;
}
*self.tail_readiness.get() = next;
tail = next;
next = (*next).next_readiness.load(Acquire);
}
if let Some(tick) = tick {
let actual = (*tail).notified_at.load(SeqCst);
if actual == tick {
return Dequeue::Yield;
}
}
if !next.is_null() {
*self.tail_readiness.get() = next;
debug_assert!(tail != self.stub());
return Dequeue::Data(tail);
}
if self.head_readiness.load(Acquire) as *const _ != tail {
return Dequeue::Inconsistent;
}
self.enqueue(self.stub());
next = (*tail).next_readiness.load(Acquire);
if !next.is_null() {
*self.tail_readiness.get() = next;
return Dequeue::Data(tail);
}
Dequeue::Inconsistent
}
fn stub(&self) -> *const Node<U> {
&*self.stub
}
}
impl<U> Drop for Inner<U> {
fn drop(&mut self) {
unsafe {
loop {
match self.dequeue(None) {
Dequeue::Empty => break,
Dequeue::Yield => unreachable!(),
Dequeue::Inconsistent => abort("inconsistent in drop"),
Dequeue::Data(ptr) => drop(ptr2arc(ptr)),
}
}
}
}
}
impl<U> List<U> {
fn new() -> Self {
List {
len: 0,
head: ptr::null_mut(),
tail: ptr::null_mut(),
}
}
fn push_back(&mut self, node: Arc<Node<U>>) -> *const Node<U> {
let ptr = arc2ptr(node);
unsafe {
*(*ptr).prev_all.get() = self.tail;
*(*ptr).next_all.get() = ptr::null_mut();
if !self.tail.is_null() {
*(*self.tail).next_all.get() = ptr;
self.tail = ptr;
} else {
self.tail = ptr;
self.head = ptr;
}
}
self.len += 1;
ptr
}
fn pop_front(&mut self) -> Option<Arc<Node<U>>> {
if self.head.is_null() {
return None;
}
self.len -= 1;
unsafe {
let node = ptr2arc(self.head);
self.head = *node.next_all.get();
if self.head.is_null() {
self.tail = ptr::null_mut();
} else {
*(*self.head).prev_all.get() = ptr::null_mut();
}
Some(node)
}
}
unsafe fn remove(&mut self, node: *const Node<U>) -> Arc<Node<U>> {
let node = ptr2arc(node);
let next = *node.next_all.get();
let prev = *node.prev_all.get();
*node.next_all.get() = ptr::null_mut();
*node.prev_all.get() = ptr::null_mut();
if !next.is_null() {
*(*next).prev_all.get() = prev;
} else {
self.tail = prev;
}
if !prev.is_null() {
*(*prev).next_all.get() = next;
} else {
self.head = next;
}
self.len -= 1;
node
}
}
unsafe fn noop(_: *const ()) {}
fn waker_inner<U: Unpark>(inner: Arc<Inner<U>>) -> Waker {
let ptr = Arc::into_raw(inner) as *const ();
let vtable = &RawWakerVTable::new(
clone_inner::<U>,
wake_inner::<U>,
wake_by_ref_inner::<U>,
drop_inner::<U>,
);
unsafe { Waker::from_raw(RawWaker::new(ptr, vtable)) }
}
unsafe fn clone_inner<U: Unpark>(data: *const ()) -> RawWaker {
let arc: Arc<Inner<U>> = Arc::from_raw(data as *const Inner<U>);
let clone = arc.clone();
mem::forget(arc);
mem::forget(clone);
let vtable = &RawWakerVTable::new(
clone_inner::<U>,
wake_inner::<U>,
wake_by_ref_inner::<U>,
drop_inner::<U>,
);
RawWaker::new(data, vtable)
}
unsafe fn wake_inner<U: Unpark>(data: *const ()) {
let arc: Arc<Inner<U>> = Arc::from_raw(data as *const Inner<U>);
arc.unpark.unpark();
}
unsafe fn wake_by_ref_inner<U: Unpark>(data: *const ()) {
let arc: Arc<Inner<U>> = Arc::from_raw(data as *const Inner<U>);
arc.unpark.unpark();
mem::forget(arc);
}
unsafe fn drop_inner<U>(data: *const ()) {
drop(Arc::<Inner<U>>::from_raw(data as *const Inner<U>));
}
unsafe fn waker_ref<U: Unpark>(node: &Arc<Node<U>>) -> Waker {
let ptr = &*node as &Node<U> as *const Node<U> as *const ();
let vtable = &RawWakerVTable::new(
clone_node::<U>,
wake_unreachable,
wake_by_ref_node::<U>,
noop,
);
Waker::from_raw(RawWaker::new(ptr, vtable))
}
unsafe fn wake_unreachable(_data: *const ()) {
unreachable!("waker_ref::wake()");
}
unsafe fn clone_node<U: Unpark>(data: *const ()) -> RawWaker {
let arc: Arc<Node<U>> = Arc::from_raw(data as *const Node<U>);
let clone = arc.clone();
mem::forget(arc);
mem::forget(clone);
let vtable = &RawWakerVTable::new(
clone_node::<U>,
wake_node::<U>,
wake_by_ref_node::<U>,
drop_node::<U>,
);
RawWaker::new(data, vtable)
}
unsafe fn wake_node<U: Unpark>(data: *const ()) {
let arc: Arc<Node<U>> = Arc::from_raw(data as *const Node<U>);
Node::<U>::notify(&arc);
}
unsafe fn wake_by_ref_node<U: Unpark>(data: *const ()) {
let arc: Arc<Node<U>> = Arc::from_raw(data as *const Node<U>);
Node::<U>::notify(&arc);
mem::forget(arc);
}
unsafe fn drop_node<U>(data: *const ()) {
drop(Arc::<Node<U>>::from_raw(data as *const Node<U>));
}
impl<U: Unpark> Node<U> {
fn notify(me: &Arc<Node<U>>) {
let inner = match me.queue.upgrade() {
Some(inner) => inner,
None => return,
};
let prev = me.queued.swap(true, SeqCst);
if !prev {
let tick_num = inner.tick_num.load(SeqCst);
me.notified_at.store(tick_num, SeqCst);
inner.enqueue(&**me);
inner.unpark.unpark();
}
}
}
impl<U> Drop for Node<U> {
fn drop(&mut self) {
unsafe {
if (*self.item.get()).is_some() {
abort("item still here when dropping");
}
}
}
}
fn arc2ptr<T>(ptr: Arc<T>) -> *const T {
let addr = &*ptr as *const T;
mem::forget(ptr);
addr
}
unsafe fn ptr2arc<T>(ptr: *const T) -> Arc<T> {
let anchor = mem::transmute::<usize, Arc<T>>(0x10);
let addr = &*anchor as *const T;
mem::forget(anchor);
let offset = addr as isize - 0x10;
mem::transmute::<isize, Arc<T>>(ptr as isize - offset)
}
fn abort(s: &str) -> ! {
struct DoublePanic;
impl Drop for DoublePanic {
fn drop(&mut self) {
panic!("panicking twice to abort the program");
}
}
let _bomb = DoublePanic;
panic!("{}", s);
}