use nanoid::nanoid;
use std::{sync::Arc, time::Duration};
use tracing::{Instrument, Level, info, span, warn};
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::{
shutdown_switch::ShutdownSwitch, stalled_to_wait_handle::stalled_to_wait,
workererror::WorkerError,
},
};
mod drop_handler;
mod lock_refresh;
mod pull_job;
pub(crate) mod shutdown_switch;
mod stalled_to_wait_handle;
mod workererror;
use pull_job::pull_job_thread;
pub struct Worker<D, R> {
uid: String,
pool: Pool,
queue_name: QueueName,
semaphore: Arc<Semaphore>,
job_receiver: Receiver<Result<JobWorkHandle<D, R>, WorkerError>>,
stalled_after: Arc<RwLock<Duration>>,
max_stalled_before_failed: Arc<RwLock<usize>>,
cooldown_after_error: Duration,
shutdown_switch: ShutdownSwitch,
join_handles: Vec<JoinHandle<()>>,
}
#[derive(Clone, Debug)]
pub struct WorkerArgs {
pub parallel_jobs: usize,
pub parallel_connections: usize,
pub max_stalled_before_failed: usize,
pub stalled_after: Duration,
pub 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 worker_id = nanoid!(8);
let worker_span = span!(
Level::ERROR,
"worker",
id = worker_id,
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;
let shutdown_switch = ShutdownSwitch::new();
let mut join_handles = Vec::new();
info!("Set up async tasks for worker");
for _ in 0..args.parallel_connections {
let pull_worker_id = nanoid!(5);
let pull_span = span!(Level::ERROR, "pull", id = pull_worker_id);
let jh = tokio::spawn(
pull_job_thread(
pool.clone(),
queue_name.clone(),
shutdown_switch.clone(),
tx.clone(),
semaphore.clone(),
args.cooldown_after_error,
pull_worker_id,
)
.instrument(pull_span.clone()),
);
join_handles.push(jh);
}
let stalled_to_wait_span = span!(Level::TRACE, "stallcheck");
join_handles.push(tokio::spawn(
stalled_to_wait(
pool.clone(),
queue_name.clone(),
shutdown_switch.clone(),
stalled_after.clone(),
max_stalled_before_failed.clone(),
)
.instrument(stalled_to_wait_span),
));
Self {
uid: worker_id,
pool,
queue_name,
semaphore,
job_receiver,
max_stalled_before_failed,
cooldown_after_error,
stalled_after,
shutdown_switch,
join_handles,
}
}
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 async fn terminate(self) {
self.shutdown_switch.shutdown();
for handle in self.join_handles {
if let Err(e) = handle.await {
warn!("Joined panicked task: {:?}", e);
}
}
}
}