use std::collections::VecDeque;
use std::future::Future;
use std::num::NonZeroUsize;
use std::pin::Pin;
use std::sync::{Arc, Mutex, MutexGuard};
use std::task::{Context, Poll, Wake, Waker};
use super::{FlowError, MAX_IN_FLIGHT, POLL_BUDGET, Scope};
pub async fn each<I, T, F, Fut>(
scope: &Scope,
step: &str,
limit: Option<NonZeroUsize>,
items: I,
body: F,
) -> Result<Vec<T>, FlowError>
where
I: IntoIterator,
F: Fn(Scope, I::Item) -> Fut,
Fut: Future<Output = Result<T, FlowError>>,
{
let here = scope.enter(step)?;
let declared = match limit {
Some(caller_limit) => caller_limit,
None => scope.policy().fan_out(),
};
let width = declared.get().min(MAX_IN_FLIGHT);
let wakes = Arc::new(Wakes::new(width));
let wakers: Vec<Waker> = (0..width)
.map(|position| {
Waker::from(Arc::new(SlotWaker {
position,
wakes: Arc::clone(&wakes),
}))
})
.collect();
let mut fan = Fan {
slots: (0..width).map(|_| None).collect(),
free: (0..width).rev().collect(),
results: Vec::new(),
exhausted: false,
};
let mut inputs = items.into_iter();
let mut stopped = Box::pin(here.token().cancelled_owned());
let body = &body;
let here = &here;
std::future::poll_fn(move |context: &mut Context<'_>| {
wakes.set_parent(context.waker());
let mut polled: usize = 0;
loop {
if stopped.as_mut().poll(context).is_ready() {
fan.clear();
return Poll::Ready(Err(FlowError::Cancelled {
at: Arc::clone(here.shared_path()),
}));
}
if let Err(error) = fan.admit(&mut inputs, here, body, &wakes) {
fan.clear();
return Poll::Ready(Err(error));
}
let ready = wakes.take_ready();
if ready.is_empty() {
break;
}
for position in ready {
polled = polled.saturating_add(1);
if let Err(error) = fan.poll_slot(position, &wakers) {
here.cancel();
fan.clear();
return Poll::Ready(Err(error));
}
}
if polled >= POLL_BUDGET {
context.waker().wake_by_ref();
return Poll::Pending;
}
}
if fan.exhausted && fan.free.len() == fan.slots.len() {
Poll::Ready(Ok(fan.results.drain(..).flatten().collect()))
} else {
Poll::Pending
}
})
.await
}
struct Slot<Fut> {
index: usize,
at: Arc<str>,
future: Pin<Box<Fut>>,
}
struct Fan<T, Fut> {
slots: Vec<Option<Slot<Fut>>>,
free: Vec<usize>,
results: Vec<Option<T>>,
exhausted: bool,
}
impl<T, Fut: Future<Output = Result<T, FlowError>>> Fan<T, Fut> {
fn admit<I, F>(
&mut self,
inputs: &mut I,
here: &Scope,
body: &F,
wakes: &Wakes,
) -> Result<(), FlowError>
where
I: Iterator,
F: Fn(Scope, I::Item) -> Fut,
{
while !self.exhausted {
let Some(position) = self.free.pop() else {
break;
};
let Some(item) = inputs.next() else {
self.free.push(position);
self.exhausted = true;
break;
};
let index = self.results.len();
self.results.push(None);
let scope = here.numbered(index)?;
let at = Arc::clone(scope.shared_path());
let future = Box::pin(body(scope, item));
if let Some(slot) = self.slots.get_mut(position) {
*slot = Some(Slot { index, at, future });
}
wakes.mark(position);
}
Ok(())
}
fn poll_slot(&mut self, position: usize, wakers: &[Waker]) -> Result<(), FlowError> {
let (Some(entry), Some(waker)) = (self.slots.get_mut(position), wakers.get(position))
else {
return Ok(());
};
let Some(slot) = entry.as_mut() else {
return Ok(());
};
let mut context = Context::from_waker(waker);
let Poll::Ready(outcome) = slot.future.as_mut().poll(&mut context) else {
return Ok(());
};
let index = slot.index;
let at = Arc::clone(&slot.at);
*entry = None;
self.free.push(position);
let value = outcome.map_err(|error| error.located_at(&at))?;
if let Some(output) = self.results.get_mut(index) {
*output = Some(value);
}
Ok(())
}
fn clear(&mut self) {
for slot in &mut self.slots {
*slot = None;
}
self.exhausted = true;
}
}
struct Wakes {
state: Mutex<WakeState>,
}
struct WakeState {
ready: VecDeque<usize>,
queued: Vec<bool>,
parent: Option<Waker>,
}
impl Wakes {
fn new(width: usize) -> Self {
Self {
state: Mutex::new(WakeState {
ready: VecDeque::with_capacity(width),
queued: vec![false; width],
parent: None,
}),
}
}
fn lock(&self) -> MutexGuard<'_, WakeState> {
crate::journal::owner::lock(&self.state)
}
fn set_parent(&self, waker: &Waker) {
let mut state = self.lock();
let stale = state
.parent
.as_ref()
.is_none_or(|current| !current.will_wake(waker));
if stale {
state.parent = Some(waker.clone());
}
}
fn mark(&self, position: usize) {
enqueue(&mut self.lock(), position);
}
fn wake_slot(&self, position: usize) -> Option<Waker> {
let mut state = self.lock();
enqueue(&mut state, position);
state.parent.clone()
}
fn take_ready(&self) -> VecDeque<usize> {
let mut state = self.lock();
let ready = std::mem::take(&mut state.ready);
for &position in &ready {
if let Some(flag) = state.queued.get_mut(position) {
*flag = false;
}
}
ready
}
}
fn enqueue(state: &mut WakeState, position: usize) {
if let Some(flag) = state.queued.get_mut(position)
&& !*flag
{
*flag = true;
state.ready.push_back(position);
}
}
struct SlotWaker {
position: usize,
wakes: Arc<Wakes>,
}
impl Wake for SlotWaker {
fn wake(self: Arc<Self>) {
self.wake_by_ref();
}
fn wake_by_ref(self: &Arc<Self>) {
if let Some(parent) = self.wakes.wake_slot(self.position) {
parent.wake();
}
}
}