use alloc::{rc::Rc, vec::Vec};
use core::{
ops::ControlFlow,
task::{Poll, Waker},
};
use crate::{
Enqueuer, Queue, Task, TaskPriority, batch::Batch, control::Control,
};
pub struct Worker<'q, T: Task> {
id: usize,
queue: &'q Queue<T>,
enqueuer: Rc<Enqueuer<T>>,
postponed: Vec<T>,
batch: Batch<T>,
control: Control,
}
impl<'q, T: Task> Worker<'q, T> {
pub fn new(queue: &'q Queue<T>, id: usize) -> Self {
assert!(id < queue.config.num_workers.get());
queue.control[id].initialize();
Self {
id,
queue,
enqueuer: Rc::new(Enqueuer::new()),
postponed: Vec::new(),
batch: unsafe { Batch::new(id, &queue.config) },
control: Control::Running,
}
}
#[inline]
pub const fn id(&self) -> usize {
self.id
}
#[inline]
pub const fn queue(&self) -> &'q Queue<T> {
self.queue
}
#[inline]
pub const fn enqueuer(&self) -> &Rc<Enqueuer<T>> {
&self.enqueuer
}
#[inline]
pub fn enqueue(&self, tasks: impl IntoIterator<Item = T>) {
self.enqueuer.extend(tasks)
}
#[inline]
pub fn enqueue_one(&self, task: T) {
self.enqueuer.add(task);
}
pub async fn next(&mut self) -> Option<T> {
core::future::poll_fn(|cx| self.poll(cx.waker())).await
}
#[inline]
pub fn poll(&mut self, waker: &Waker) -> Poll<Option<T>> {
if let Some(task) = unsafe { self.batch.next(self.id, self.queue) } {
debug_assert_eq!(self.control, Control::Running);
return Poll::Ready(Some(task));
}
self.poll_slow(waker)
}
#[cold]
fn poll_slow(&mut self, waker: &Waker) -> Poll<Option<T>> {
if self.control == Control::Running {
self.refresh_and_poll(waker)
} else if self.control == Control::Asleep {
let control = &self.queue.control[self.id];
let waker_slot = &self.queue.waker_slot[self.id];
if self.enqueuer.take_waker().is_none() {
unsafe { control.set_self_awake(waker_slot) };
self.control = Control::Running;
} else {
self.control =
unsafe { control.poll_asleep(waker, waker_slot) };
}
if matches!(self.control, Control::Asleep) {
self.enqueuer.set_waker(waker.clone());
Poll::Pending
} else {
if self.queue.queue_control.mark_awake().is_break() {
self.control = Control::ShuttingDown;
return Poll::Ready(None);
}
self.poll(waker)
}
} else if self.control == Control::Initialized {
self.control = Control::Running;
self.poll(waker)
} else {
Poll::Ready(None)
}
}
fn refresh_and_poll(&mut self, waker: &Waker) -> Poll<Option<T>> {
if self.queue.queue_control.has_shut_down() {
return Poll::Ready(None);
}
self.collect();
let task = self.postponed.pop();
unsafe {
self.batch.fill(
task.is_some() as usize,
&mut self.postponed,
self.id,
self.queue,
)
};
match task {
Some(task) => Poll::Ready(Some(task)),
None => {
let control = &self.queue.control[self.id];
let waker_slot = &self.queue.waker_slot[self.id];
unsafe { control.set_self_asleep(waker_slot, waker.clone()) };
self.enqueuer.set_waker(waker.clone());
self.control = match self.queue.queue_control.mark_asleep() {
ControlFlow::Continue(0) if self.queue.config.oneshot => {
self.queue.shutdown();
Control::ShuttingDown
}
ControlFlow::Continue(_) => Control::Asleep,
ControlFlow::Break(_) => Control::ShuttingDown,
};
match self.control {
Control::Asleep => Poll::Pending,
Control::ShuttingDown => Poll::Ready(None),
_ => unreachable!(),
}
}
}
}
fn collect(&mut self) {
unsafe { self.batch.drain(&mut self.postponed, self.id, self.queue) };
let max_priority = self.postponed.last().map(|t| t.priority());
self.enqueuer.drain(&mut self.postponed);
let pq = self.batch.pub_queue.id(&self.queue.config);
if let Some(pq) = unsafe { self.queue.steal(max_priority, pq) } {
let stealer = &self.queue.stealer[self.id];
unsafe { stealer.set(self.batch.pub_queue, pq) };
self.batch.pub_queue = pq;
let pq = self.batch.pub_queue.id(&self.queue.config);
let len = self.batch.pub_queue.len(&self.queue.config);
let pq = &self.queue.pq_contents[pq];
unsafe { pq.move_to(len.get(), &mut self.postponed) };
}
let batch_size = self.queue.config.batch_size.get();
if <T::Priority as TaskPriority>::TRIVIAL {
} else if self.postponed.len() > batch_size * 2 {
let sort_offset = self.postponed.len() - batch_size * 2;
self.postponed
.select_nth_unstable_by_key(sort_offset, |task| {
task.priority()
});
self.postponed[sort_offset..]
.sort_unstable_by_key(|task| task.priority());
} else {
self.postponed.sort_unstable_by_key(|task| task.priority());
}
debug_assert!(self.batch.counter.is_empty());
debug_assert!(self.batch.local.is_empty());
debug_assert!(self.batch.pq_priority.is_none());
}
}
impl<T: Task> Drop for Worker<'_, T> {
fn drop(&mut self) {
unsafe { self.batch.drain(&mut self.postponed, self.id, self.queue) };
}
}
#[cfg(test)]
mod test {
use core::{
convert::Infallible,
iter,
num::{NonZeroU8, NonZeroUsize},
task::{Poll, Waker},
};
use crate::{Config, Task, Worker};
#[test]
fn new() {
let workers = 8.try_into().unwrap();
let batch_size = NonZeroUsize::MIN;
let queue = Config::new(workers)
.with_batch_size(batch_size)
.build::<Infallible>();
for id in 0..8 {
let _ = Worker::new(&queue, id);
}
}
#[test]
#[should_panic]
fn new_too_many() {
let workers = 4.try_into().unwrap();
let batch_size = NonZeroUsize::MIN;
let queue = Config::new(workers)
.with_batch_size(batch_size)
.build::<Infallible>();
for id in 0..5 {
let _ = Worker::new(&queue, id);
}
}
#[test]
fn next_single() {
struct Foo(u8);
impl Task for Foo {
type Priority = ();
fn priority(&self) -> Self::Priority {}
}
let workers = 1.try_into().unwrap();
let batch_size = NonZeroUsize::MIN;
let queue = Config::new(workers).with_batch_size(batch_size).build();
let mut worker = Worker::new(&queue, 0);
worker.enqueuer().extend((0..4).map(Foo));
let mut seen = [false; 4];
let waker = Waker::noop();
while let Poll::Ready(Some(task)) = worker.poll(waker) {
let Foo(index) = task;
assert!(!seen[index as usize]);
seen[index as usize] = true;
}
assert_eq!(seen, [true; 4]);
}
#[test]
fn priority() {
struct Foo(NonZeroU8);
impl Task for Foo {
type Priority = NonZeroU8;
fn priority(&self) -> Self::Priority {
self.0
}
}
let workers = 1.try_into().unwrap();
let batch_size = NonZeroUsize::MIN;
let queue = Config::new(workers).with_batch_size(batch_size).build();
let mut worker = Worker::new(&queue, 0);
worker
.enqueuer()
.extend((1..5).rev().map(|i| Foo(NonZeroU8::new(i).unwrap())));
let waker = Waker::noop();
assert!(
iter::from_fn(|| match worker.poll(waker) {
Poll::Ready(Some(t)) => Some(t),
_ => None,
})
.map(|t| t.0.get())
.eq((1..5).rev())
);
}
}