use std::future::Future;
use std::io;
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::pin::Pin;
use std::sync::mpsc::{self, Receiver, SyncSender, TrySendError};
use std::sync::{Arc, Mutex, PoisonError};
use std::task::{Context, Poll};
use moirai_pal::thread::ThreadStartError;
use crate::sync::{Semaphore, SemaphorePermit, oneshot};
#[cfg(test)]
pub(crate) mod job_lifecycle;
#[cfg(test)]
pub(crate) mod test_hooks;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum Abandoned {
Skip,
Run,
}
type Job = Box<dyn FnOnce() + Send>;
pub(crate) struct BlockingPool {
name: &'static str,
workers_bound: usize,
#[cfg(test)]
admissions: usize,
admission: Semaphore,
jobs: SyncSender<Job>,
queue: Arc<Mutex<Receiver<Job>>>,
workers: Mutex<usize>,
#[cfg(test)]
hooks: test_hooks::Hooks,
}
impl BlockingPool {
pub(crate) fn new(name: &'static str, workers: usize, queue_depth: usize) -> Self {
let admissions = workers + queue_depth;
let (jobs, queue) = mpsc::sync_channel(admissions);
Self {
name,
workers_bound: workers,
#[cfg(test)]
admissions,
admission: Semaphore::new(admissions),
jobs,
queue: Arc::new(Mutex::new(queue)),
workers: Mutex::new(0),
#[cfg(test)]
hooks: test_hooks::Hooks::new(),
}
}
pub(crate) async fn admit(&'static self) -> io::Result<Admission> {
self.ensure_workers()?;
let permit = self.admission.acquire().await;
Ok(Admission { pool: self, permit })
}
pub(crate) async fn run<T, F>(&'static self, abandoned: Abandoned, work: F) -> io::Result<T>
where
T: Send + 'static,
F: FnOnce() -> T + Send + 'static,
{
self.admit().await?.submit(abandoned, work)?.await
}
fn ensure_workers(&'static self) -> io::Result<()> {
let mut workers = self.workers.lock().unwrap_or_else(PoisonError::into_inner);
while *workers < self.workers_bound {
let spawned = std::thread::Builder::new()
.name(format!("{}-{}", self.name, *workers))
.spawn(move || self.serve());
match spawned {
Ok(_) => *workers += 1,
Err(source) if *workers == 0 => {
return Err(ThreadStartError::new(self.name, source).into());
}
Err(_) => break,
}
}
Ok(())
}
fn serve(&self) {
loop {
let next = self
.queue
.lock()
.unwrap_or_else(PoisonError::into_inner)
.recv();
let Ok(job) = next else {
return;
};
let _contained = catch_unwind(AssertUnwindSafe(job));
#[cfg(test)]
self.hooks.disposed();
}
}
}
pub(crate) struct Admission {
pool: &'static BlockingPool,
permit: SemaphorePermit<'static>,
}
impl Admission {
pub(crate) fn submit<T, F>(self, abandoned: Abandoned, work: F) -> io::Result<Completion<T>>
where
T: Send + 'static,
F: FnOnce() -> T + Send + 'static,
{
let Self { pool, permit } = self;
let (reply, receiver) = oneshot::channel();
let job: Job = Box::new(move || {
if abandoned == Abandoned::Run || !reply.is_closed() {
#[cfg(test)]
let _running = pool.hooks.enter();
let result = work();
drop(reply.send(result));
}
drop(permit);
});
match pool.jobs.try_send(job) {
Ok(()) => Ok(Completion {
pool: pool.name,
receiver,
}),
Err(TrySendError::Full(_)) => unreachable!(
"invariant: the channel holds one slot per admission permit, so an \
admitted job always finds a queue slot"
),
Err(TrySendError::Disconnected(_)) => Err(io::Error::other(format!(
"the {} queue closed before the job was submitted",
pool.name
))),
}
}
}
pub(crate) struct Completion<T> {
pool: &'static str,
receiver: oneshot::Receiver<T>,
}
impl<T> Future for Completion<T> {
type Output = io::Result<T>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let pool = self.pool;
self.receiver.poll_recv(cx).map(|received| {
received
.map_err(|()| io::Error::other(format!("a {pool} job panicked before replying")))
})
}
}
#[cfg(test)]
impl BlockingPool {
pub(crate) fn hooks(&self) -> &test_hooks::Hooks {
&self.hooks
}
pub(crate) fn admissions(&self) -> usize {
self.admissions
}
pub(crate) fn workers_bound(&self) -> usize {
self.workers_bound
}
pub(crate) fn workers(&self) -> usize {
*self.workers.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn free_admissions(&self) -> usize {
self.admission.available_permits()
}
}