use alloc::vec::Vec;
use core::{
sync::atomic::Ordering,
task::{Poll, Waker},
};
use std::rc::Rc;
use crate::{Enqueuer, Queue, Task, 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>,
executed: usize,
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());
let control = match queue.control[id].initialize() {
false => Control::Initialized,
true => Control::ShuttingDown,
};
Self {
id,
queue,
enqueuer: Rc::new(Enqueuer::new()),
postponed: Vec::new(),
batch: unsafe { Batch::new(id, &queue.config) },
executed: 0,
control,
}
}
#[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);
self.executed += 1;
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 {
let control = &self.queue.control[self.id];
self.control = unsafe { control.poll_running() };
if matches!(self.control, Control::Running) {
self.refresh_and_poll(waker)
} else {
self.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() {
self.control = unsafe { control.set_self_awake(waker_slot) };
} 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 {
self.poll(waker)
}
} else if self.control == Control::Initialized {
self.control = Control::Running;
self.queue.started.fetch_add(1, Ordering::Relaxed);
self.poll(waker)
} else {
Poll::Ready(None)
}
}
fn refresh_and_poll(&mut self, waker: &Waker) -> Poll<Option<T>> {
self.collect();
let executed = core::mem::take(&mut self.executed);
let empty = if self.queue.config.oneshot {
self.queue.pending.fetch_sub(executed, Ordering::Relaxed)
<= executed
} else {
false
};
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) => {
self.executed += 1;
Poll::Ready(Some(task))
}
None => {
if empty && self.queue.all_started() {
self.queue.shutdown();
self.control = Control::ShuttingDown;
return Poll::Ready(None);
}
let control = &self.queue.control[self.id];
let waker_slot = &self.queue.waker_slot[self.id];
self.control = unsafe {
control.set_self_asleep(waker_slot, waker.clone())
};
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());
let enqueued = self.enqueuer.drain(&mut self.postponed);
if self.queue.config.oneshot {
self.queue.pending.fetch_add(enqueued, Ordering::Relaxed);
}
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) };
}
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())
);
}
}