use nanoid::nanoid;
use std::{
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
time::Duration,
};
use tracing::{Instrument, Level, info, span};
use deadpool_redis::Pool;
use serde::de::DeserializeOwned;
use tokio::{
sync::{
RwLock, Semaphore,
mpsc::{Receiver, channel},
},
task::JoinHandle,
};
use crate::{
job::JobWorkHandle,
queue::QueueName,
worker::{stalled_to_wait_handle::stalled_to_wait, workererror::WorkerError},
};
mod drop_handler;
mod pull_job;
mod stalled_to_wait_handle;
mod workererror;
use pull_job::pull_job_thread;
pub struct Worker<D, R> {
pool: Pool,
queue_name: QueueName,
semaphore: Arc<Semaphore>,
job_fetch_handles: Vec<JoinHandle<()>>,
stalled_to_wait_handle: JoinHandle<()>,
job_receiver: Receiver<Result<JobWorkHandle<D, R>, WorkerError>>,
stalled_after: Arc<RwLock<Duration>>,
max_stalled_before_failed: Arc<RwLock<usize>>,
cooldown_after_error: Duration,
uid: String,
terminating_initiated: AtomicBool,
}
#[derive(Clone, Debug)]
pub struct WorkerArgs {
pub parallel_jobs: usize,
pub parallel_connections: usize,
pub max_stalled_before_failed: usize,
pub stalled_after: Duration,
cooldown_after_error: Duration,
}
impl Default for WorkerArgs {
fn default() -> Self {
Self {
parallel_jobs: 32,
parallel_connections: 1,
max_stalled_before_failed: 1,
stalled_after: Duration::from_secs(30),
cooldown_after_error: Duration::from_secs(3),
}
}
}
impl<D, R> Worker<D, R>
where
R: Send + 'static,
D: Send + Sync + 'static + DeserializeOwned + std::fmt::Debug,
{
pub(crate) fn new(pool: Pool, queue_name: QueueName, args: WorkerArgs) -> Self {
let uid = nanoid!(8);
let worker_span = span!(
Level::ERROR,
"worker",
worker = uid,
queue = queue_name.as_str()
);
let _guard = worker_span.enter();
let semaphore = Arc::new(Semaphore::new(args.parallel_jobs));
let (tx, job_receiver) = channel(args.parallel_jobs);
let stalled_after = Arc::new(RwLock::new(args.stalled_after));
let max_stalled_before_failed = Arc::new(RwLock::new(args.max_stalled_before_failed));
let cooldown_after_error = args.cooldown_after_error;
info!("Set up async tasks for worker");
let job_fetch_handles: Vec<_> = (0..args.parallel_connections)
.map(|_| {
let pull_worker_id = nanoid!(5);
let pull_span = span!(Level::ERROR, "pull-jobs", pull_worker_id);
tokio::spawn(
pull_job_thread(
pool.clone(),
queue_name.clone(),
tx.clone(),
semaphore.clone(),
args.cooldown_after_error,
pull_worker_id,
)
.instrument(pull_span.clone()),
)
})
.collect();
let stalled_to_wait_span = span!(Level::TRACE, "stalled-to-wait");
let stalled_to_wait_handle = tokio::spawn(
stalled_to_wait(
pool.clone(),
queue_name.clone(),
stalled_after.clone(),
max_stalled_before_failed.clone(),
)
.instrument(stalled_to_wait_span),
);
Self {
uid,
pool,
queue_name,
semaphore,
job_fetch_handles,
job_receiver,
max_stalled_before_failed,
stalled_to_wait_handle,
cooldown_after_error,
stalled_after,
terminating_initiated: AtomicBool::from(false),
}
}
pub async fn next(&mut self) -> Option<Result<JobWorkHandle<D, R>, WorkerError>> {
self.job_receiver.recv().await
}
pub fn has_next(&self) -> bool {
!self.job_receiver.is_empty()
}
pub fn terminate(&self) {
self.terminating_initiated.store(true, Ordering::SeqCst);
self.stalled_to_wait_handle.abort();
for h in self.job_fetch_handles.iter() {
h.abort();
}
}
fn is_terminating_gracefully(&self) -> bool {
self.terminating_initiated.load(Ordering::SeqCst)
}
}
impl<D, R> Drop for Worker<D, R> {
fn drop(&mut self) {
self.stalled_to_wait_handle.abort();
self.job_fetch_handles.iter().for_each(|h| h.abort());
}
}