#[cfg(feature = "prometheus")]
use crate::metrics;
use crate::{
commands,
core::BgJobHandler,
encoder,
models::{AmqpCommand, ChannelCommand},
mq::MqClient,
BackgroundJobServer, BackgroundJobServerPublisher, ServerConfig, UtcDateTime,
};
use async_std::channel::{Receiver, Sender};
use std::{marker::PhantomData, sync::Arc, time::Duration};
pub(crate) struct WorkerTask {
stop: Option<tokio::sync::oneshot::Sender<()>>,
task: tokio::task::JoinHandle<anyhow::Result<()>>,
}
impl<C, H> BackgroundJobServer<C, H>
where
C: Sync + Send + 'static,
H: BgJobHandler<C> + Sync + Send + 'static,
{
pub async fn start(
handler: H,
mq_client: Arc<Box<dyn MqClient>>,
config: ServerConfig,
) -> anyhow::Result<Self> {
let mut worker_tasks = Vec::new();
let mut maintenance_tasks = Vec::new();
let mut partition_membership_task = None;
let mut partition_tasks = Vec::new();
let mut partition_stop = None;
let handler = Arc::new(handler);
let publisher = handler.get_publisher();
let num_bg_workers = config.worker_count;
if num_bg_workers == 0 {
return Err(anyhow::anyhow!("worker_count must be at least 1"));
}
let (tx, rx) = async_std::channel::unbounded::<ChannelCommand>();
let (worker_ready_tx, worker_ready_rx) =
async_std::channel::bounded::<()>(num_bg_workers.into());
#[cfg(feature = "dashboard")]
let server_instance_id = crate::generate_id();
for id in 0..num_bg_workers {
worker_tasks.push(spawn_worker(
handler.clone(),
mq_client.clone(),
publisher.routing_key.clone(),
tx.clone(),
id.into(),
worker_ready_tx.clone(),
#[cfg(feature = "dashboard")]
server_instance_id.clone(),
));
}
drop(worker_ready_tx);
for _ in 0..num_bg_workers {
if worker_ready_rx.recv().await.is_err() {
for worker in &worker_tasks {
worker.task.abort();
}
return Err(anyhow::anyhow!("job worker stopped during startup"));
}
}
let handler_for_ensure_ops = handler.clone();
let ensure_ops_tx = tx.clone();
maintenance_tasks.push(tokio::spawn(async move {
start_bg_worker_system_ops_ensure_ops_run_on_certain_interval(
handler_for_ensure_ops,
ensure_ops_tx,
)
.await
}));
#[cfg(feature = "prometheus")]
{
let metrics_handler = handler.clone();
maintenance_tasks.push(tokio::spawn(async move {
publish_queue_depths(metrics_handler).await
}));
}
let mut partition_owner = None;
if handler.get_publisher().has_topics() {
let owner = crate::generate_id();
handler
.get_publisher()
.heartbeat_partition_worker(&owner)
.await?;
let live_workers = std::sync::Arc::new(tokio::sync::Mutex::new(vec![owner.clone()]));
let (stop_tx, _) = tokio::sync::watch::channel(false);
{
let membership_handler = handler.clone();
let membership_owner = owner.clone();
let membership_live_workers = live_workers.clone();
partition_membership_task = Some(tokio::spawn(async move {
run_partition_membership(
membership_handler,
membership_owner,
membership_live_workers,
)
.await
}));
}
for _ in 0..num_bg_workers {
let partition_handler = handler.clone();
let owner = owner.clone();
let live_workers = live_workers.clone();
let stop = stop_tx.subscribe();
partition_tasks.push(tokio::spawn(async move {
run_partition_poller(partition_handler, owner, live_workers, stop).await
}));
}
partition_owner = Some(owner);
partition_stop = Some(stop_tx);
}
for _ in 1..2 {
let rx_clone = rx.clone();
let tx_clone = tx.clone();
let handler_for_check_ops = handler.clone();
maintenance_tasks.push(tokio::spawn(async move {
start_bg_worker_system_ops_inproc_cmd_to_amqp_cmd(
handler_for_check_ops,
tx_clone,
rx_clone,
)
.await
}));
}
Ok(Self {
ctx: PhantomData,
handler,
mq_client,
inproc_cmd_tx: tx,
worker_tasks: std::sync::Mutex::new(worker_tasks),
worker_scaling: tokio::sync::Mutex::new(()),
next_worker_id: std::sync::atomic::AtomicU32::new(num_bg_workers.into()),
#[cfg(feature = "dashboard")]
server_instance_id,
maintenance_tasks,
partition_owner,
partition_membership_task,
partition_tasks: std::sync::Mutex::new(partition_tasks),
partition_stop,
})
}
pub fn worker_count(&self) -> usize {
self.with_worker_tasks(|workers| {
workers
.iter()
.filter(|worker| !worker.task.is_finished())
.count()
})
}
pub async fn add_worker(&self) -> anyhow::Result<usize> {
let _scaling = self.worker_scaling.lock().await;
let current_count = self.with_worker_tasks(|workers| {
workers.retain(|worker| !worker.task.is_finished());
workers.len()
});
if current_count >= usize::from(u8::MAX) {
return Err(anyhow::anyhow!("worker count cannot exceed 255"));
}
let worker_id = self
.next_worker_id
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let worker_id = i32::try_from(worker_id)
.map_err(|_| anyhow::anyhow!("worker identifier is out of range"))?;
let (worker_ready_tx, worker_ready_rx) = async_std::channel::bounded(1);
let worker = spawn_worker(
self.handler.clone(),
self.mq_client.clone(),
self.handler.get_publisher().routing_key.clone(),
self.inproc_cmd_tx.clone(),
worker_id,
worker_ready_tx,
#[cfg(feature = "dashboard")]
self.server_instance_id.clone(),
);
if worker_ready_rx.recv().await.is_err() {
worker.task.abort();
return Err(anyhow::anyhow!("worker stopped during startup"));
}
Ok(self.with_worker_tasks(|workers| {
workers.push(worker);
workers.len()
}))
}
pub async fn remove_worker(&self) -> anyhow::Result<usize> {
let _scaling = self.worker_scaling.lock().await;
let (worker, remaining) = self
.with_worker_tasks(|workers| {
workers.retain(|worker| !worker.task.is_finished());
if workers.len() <= 1 {
return None;
}
let worker = workers.pop()?;
Some((worker, workers.len()))
})
.ok_or_else(|| anyhow::anyhow!("at least one worker must remain"))?;
let WorkerTask { stop, task } = worker;
if let Some(stop) = stop {
let _ = stop.send(());
}
task.await
.map_err(|error| anyhow::anyhow!("worker task stopped unexpectedly: {error}"))??;
Ok(remaining)
}
pub async fn shutdown(mut self, timeout: Duration) -> anyhow::Result<()> {
for task in &self.maintenance_tasks {
task.abort();
}
self.maintenance_tasks.clear();
if let Some(stop) = self.partition_stop.take() {
let _ = stop.send(true);
}
let mut workers = self.with_worker_tasks(std::mem::take);
for worker in &mut workers {
if let Some(stop) = worker.stop.take() {
let _ = stop.send(());
}
}
let mut partition_tasks = match self.partition_tasks.lock() {
Ok(mut tasks) => std::mem::take(&mut *tasks),
Err(poisoned) => std::mem::take(&mut *poisoned.into_inner()),
};
let drained = async {
for worker in &mut workers {
(&mut worker.task).await.map_err(|error| {
anyhow::anyhow!("worker task stopped unexpectedly: {error}")
})??;
}
for task in &mut partition_tasks {
task.await.map_err(|error| {
anyhow::anyhow!("partition poller stopped unexpectedly: {error}")
})??;
}
Ok::<(), anyhow::Error>(())
};
let result = match tokio::time::timeout(timeout, drained).await {
Ok(result) => result,
Err(_) => {
for worker in &workers {
worker.task.abort();
}
for task in &partition_tasks {
task.abort();
}
Err(anyhow::anyhow!(
"graceful shutdown timed out; durable jobs will recover through their leases"
))
}
};
if let Some(task) = self.partition_membership_task.take() {
task.abort();
}
if let Some(owner) = self.partition_owner.take() {
self.handler
.get_publisher()
.deregister_partition_worker(&owner)
.await?;
}
result
}
fn with_worker_tasks<T>(&self, operation: impl FnOnce(&mut Vec<WorkerTask>) -> T) -> T {
match self.worker_tasks.lock() {
Ok(mut workers) => operation(&mut workers),
Err(poisoned) => {
let mut workers = poisoned.into_inner();
operation(&mut workers)
}
}
}
}
#[cfg(feature = "dashboard")]
async fn publish_worker_heartbeats<C, H>(handler: Arc<H>, worker_id: String) -> anyhow::Result<()>
where
C: Sync + Send,
H: BgJobHandler<C> + Sync + Send + 'static,
{
let mut interval = tokio::time::interval(Duration::from_secs(5));
loop {
interval.tick().await;
if let Err(error) = handler
.get_publisher()
.publish_worker_heartbeat(worker_id.clone())
.await
{
tracing::warn!(%worker_id, %error, "Could not publish worker heartbeat");
}
}
}
#[cfg(feature = "prometheus")]
async fn publish_queue_depths<C, H>(handler: Arc<H>) -> anyhow::Result<()>
where
C: Sync + Send,
H: BgJobHandler<C> + Sync + Send + 'static,
{
let mut interval = tokio::time::interval(Duration::from_secs(5));
loop {
interval.tick().await;
match handler.get_publisher().storage.queue_depths().await {
Ok(depths) => {
for depth in depths {
metrics::set_queue_depth(
handler.get_publisher().metrics_queue(),
depth.state.metric_name(),
depth.count,
);
}
}
Err(error) => tracing::warn!(%error, "Could not collect queue depth metrics"),
}
}
}
async fn run_partition_membership<C, H>(
handler: Arc<H>,
owner: String,
live_workers: Arc<tokio::sync::Mutex<Vec<String>>>,
) -> anyhow::Result<()>
where
C: Sync + Send,
H: BgJobHandler<C> + Sync + Send + 'static,
{
loop {
if let Err(error) = handler
.get_publisher()
.heartbeat_partition_worker(&owner)
.await
{
tracing::warn!(%error, "Failed to send partition worker heartbeat");
}
match handler.get_publisher().list_live_partition_workers().await {
Ok(mut workers) => {
if !workers.iter().any(|live| live == &owner) {
workers.push(owner.clone());
}
#[cfg(feature = "prometheus")]
metrics::set_partition_workers_live(
handler.get_publisher().metrics_queue(),
workers.len(),
);
*live_workers.lock().await = workers;
}
Err(error) => tracing::warn!(%error, "Failed to list live partition workers"),
}
#[cfg(feature = "prometheus")]
match handler.get_publisher().partition_backlog_summary().await {
Ok(summary) => {
let mut by_topic: std::collections::HashMap<String, (usize, f64)> = summary
.into_iter()
.map(|(topic, ready_count, age)| (topic, (ready_count as usize, age)))
.collect();
for topic in handler.get_publisher().topic_names() {
let (ready_count, oldest_head_age_seconds) =
by_topic.remove(topic).unwrap_or((0, 0.0));
metrics::set_partition_backlog(
handler.get_publisher().metrics_queue(),
topic,
ready_count,
oldest_head_age_seconds,
);
}
}
Err(error) => tracing::warn!(%error, "Failed to summarize partition backlog"),
}
#[cfg(feature = "prometheus")]
match handler
.get_publisher()
.partition_queue_depth_summary()
.await
{
Ok(summary) => {
let mut by_partition: std::collections::HashMap<(String, u32), u64> = summary
.into_iter()
.map(|(topic, partition, depth)| ((topic, partition.0), depth))
.collect();
for topic_config in handler.get_publisher().topics() {
for partition in 0..topic_config.partition_count() {
let depth = by_partition
.remove(&(topic_config.name().to_string(), partition))
.unwrap_or(0);
metrics::set_partition_queue_depth(
handler.get_publisher().metrics_queue(),
topic_config.name(),
&partition.to_string(),
i64::try_from(depth).unwrap_or(i64::MAX),
);
}
}
}
Err(error) => tracing::warn!(%error, "Failed to summarize partition queue depth"),
}
#[cfg(all(feature = "prometheus", feature = "retained-log"))]
if let Some(source) = handler.get_publisher().retained_log_lag().source() {
for topic in handler.get_publisher().retained_log_lag().topics() {
match source.consumer_lag(topic).await {
Ok(lags) => {
for lag in lags {
metrics::set_retained_log_consumer_lag(
handler.get_publisher().metrics_queue(),
topic,
&lag.group,
lag.lag,
);
}
}
Err(error) => {
tracing::warn!(%error, topic, "Failed to compute consumer lag")
}
}
}
}
sleep_ms(2_000).await;
}
}
async fn run_partition_poller<C, H>(
handler: Arc<H>,
owner: String,
live_workers: Arc<tokio::sync::Mutex<Vec<String>>>,
mut stop: tokio::sync::watch::Receiver<bool>,
) -> anyhow::Result<()>
where
C: Sync + Send,
H: BgJobHandler<C> + Sync + Send + 'static,
{
#[cfg(feature = "prometheus")]
let worker_id_label = owner.clone();
let wake = handler.get_publisher().partition_wake().clone();
let mut assignment: Option<crate::partition::PartitionAssignment> = None;
loop {
if *stop.borrow() {
return Ok(());
}
let workers = live_workers.lock().await.clone();
match crate::partition::poll_partition_head(
&handler,
&owner,
&workers,
&mut assignment,
#[cfg(feature = "prometheus")]
&worker_id_label,
)
.await
{
Ok(true) => continue,
Ok(false) => {
tokio::select! {
_ = sleep_ms(250) => {}
_ = wake.notified() => {}
changed = stop.changed() => {
if changed.is_ok() && *stop.borrow() {
return Ok(());
}
}
}
}
Err(error) => {
tracing::warn!(%error, "Partition head poll failed");
sleep_ms(1_000).await;
}
}
}
}
#[allow(clippy::too_many_arguments)]
fn spawn_worker<C, H>(
handler: Arc<H>,
mq_client: Arc<Box<dyn MqClient>>,
routing_key: String,
inproc_cmd_tx: Sender<ChannelCommand>,
worker_id: i32,
worker_ready_tx: Sender<()>,
#[cfg(feature = "dashboard")] server_instance_id: String,
) -> WorkerTask
where
C: Sync + Send + 'static,
H: BgJobHandler<C> + Sync + Send + 'static,
{
let (stop, stop_rx) = tokio::sync::oneshot::channel();
let task = tokio::spawn(async move {
#[cfg(feature = "dashboard")]
let heartbeat = {
let heartbeat_handler = handler.clone();
let heartbeat_worker_id = format!("{server_instance_id}:{worker_id}");
tokio::spawn(async move {
publish_worker_heartbeats(heartbeat_handler, heartbeat_worker_id).await
})
};
let result = start_distributed_job_worker(
handler,
worker_id,
mq_client,
&routing_key,
inproc_cmd_tx,
worker_ready_tx,
stop_rx,
)
.await;
#[cfg(feature = "dashboard")]
heartbeat.abort();
result
});
WorkerTask {
stop: Some(stop),
task,
}
}
pub(crate) async fn sleep_ms(ms: u64) {
tokio::time::sleep(Duration::from_millis(ms)).await;
}
impl<C, H> std::ops::Deref for BackgroundJobServer<C, H>
where
C: Sync + Send + 'static,
H: BgJobHandler<C> + Sync + Send + 'static,
{
type Target = BackgroundJobServerPublisher;
fn deref(&self) -> &Self::Target {
self.handler.get_publisher()
}
}
async fn start_bg_worker_system_ops_inproc_cmd_to_amqp_cmd<C, H>(
handler: Arc<H>,
tx: Sender<ChannelCommand>,
rx: Receiver<ChannelCommand>,
) -> anyhow::Result<()>
where
C: Sync + Send,
H: BgJobHandler<C> + Sync + Send + 'static,
{
loop {
let channel_command = match rx.recv().await {
Ok(command) => command,
Err(_) => {
tracing::info!("System operation channel closed");
return Ok(());
}
};
sleep_ms(2000).await;
let config = handler.get_publisher().storage.config();
match channel_command {
ChannelCommand::PollDelayedJobs => {
commands::handle_poll_delayed_job_command(handler.clone()).await?;
config.poll_delayed_jobs_last_run_set().await?;
tx.send(ChannelCommand::PollDelayedJobs).await?;
}
ChannelCommand::PollRequeuedJobs => {
commands::handle_poll_requeued_job_command(handler.clone()).await?;
config.poll_requeued_jobs_last_run_set().await?;
tx.send(ChannelCommand::PollRequeuedJobs).await?;
}
ChannelCommand::PollRecurringJobs => {
commands::handle_poll_recurring_job_command(handler.clone()).await?;
config.poll_recurring_jobs_last_run_set().await?;
tx.send(ChannelCommand::PollRecurringJobs).await?;
}
ChannelCommand::PollExpiredStorage => {
commands::handle_poll_expired_storage_command(handler.clone()).await?;
config.poll_expired_storage_last_run_set().await?;
tx.send(ChannelCommand::PollExpiredStorage).await?;
}
ChannelCommand::PollStuckJobs => {
commands::handle_poll_stuck_jobs_command(handler.clone()).await?;
config.poll_stuck_jobs_last_run_set().await?;
tx.send(ChannelCommand::PollStuckJobs).await?;
}
}
}
}
async fn start_bg_worker_system_ops_ensure_ops_run_on_certain_interval<C, H>(
handler: Arc<H>,
tx: Sender<ChannelCommand>,
) -> anyhow::Result<()>
where
C: Sync + Send,
H: BgJobHandler<C> + Sync + Send + 'static,
{
sleep_ms(3_000).await;
loop {
let config = handler.get_publisher().storage.config();
enqueue_if(
&tx,
ChannelCommand::PollDelayedJobs,
config.poll_delayed_jobs_last_run().await?,
10,
)
.await?;
enqueue_if(
&tx,
ChannelCommand::PollRequeuedJobs,
config.poll_requeued_jobs_last_run().await?,
10,
)
.await?;
enqueue_if(
&tx,
ChannelCommand::PollRecurringJobs,
config.poll_recurring_jobs_last_run().await?,
10,
)
.await?;
enqueue_if(
&tx,
ChannelCommand::PollExpiredStorage,
config.poll_expired_storage_last_run().await?,
15,
)
.await?;
enqueue_if(
&tx,
ChannelCommand::PollStuckJobs,
config.poll_stuck_jobs_last_run().await?,
30,
)
.await?;
sleep_ms(10_000).await; }
async fn enqueue_if(
tx: &Sender<ChannelCommand>,
cmd: ChannelCommand,
last_run: UtcDateTime,
if_not_sec: i64,
) -> anyhow::Result<()> {
let last_run_since = chrono::Utc::now() - last_run;
if last_run_since.num_seconds() > if_not_sec {
tracing::warn!("Ops {} did not run for a while: Enqueuing", cmd);
tx.send(cmd).await?;
}
Ok(())
}
}
async fn start_distributed_job_worker<C, H>(
handler: Arc<H>,
worker_id: i32,
mq_client: Arc<Box<dyn MqClient>>,
routing_key: &str,
inproc_cmd_tx: Sender<ChannelCommand>,
worker_ready_tx: Sender<()>,
mut stop: tokio::sync::oneshot::Receiver<()>,
) -> anyhow::Result<()>
where
C: Sync + Send,
H: BgJobHandler<C> + Sync + Send + 'static,
{
tracing::info!("[Worker#{}] Starting", worker_id);
#[cfg(feature = "prometheus")]
let worker_id_label = worker_id.to_string();
#[cfg(feature = "prometheus")]
let _worker_guard = metrics::WorkerGuard::new(routing_key, &worker_id_label);
let mut consumer = mq_client.new_consumer(routing_key, worker_id).await?;
worker_ready_tx
.send(())
.await
.map_err(|_| anyhow::anyhow!("server stopped during worker startup"))?;
loop {
let message = tokio::select! {
biased;
_ = &mut stop => return Ok(()),
message = consumer.next() => message,
};
let Some(message) = message else {
return Ok(());
};
match message {
Ok(delivery) => match encoder::decode::<AmqpCommand>(delivery.data()) {
Ok(command) => {
let headers = delivery.get_headers();
#[cfg(feature = "prometheus")]
let command_name = metrics::command_name(&command);
match commands::handle_amqp_command(
command,
worker_id,
#[cfg(feature = "prometheus")]
&worker_id_label,
&handler,
&inproc_cmd_tx,
headers,
)
.await
{
Ok(()) => {
#[cfg(feature = "prometheus")]
metrics::record_command(
routing_key,
&worker_id_label,
command_name,
"success",
);
delivery.ack().await?
}
Err(error) => {
#[cfg(feature = "prometheus")]
metrics::record_command(
routing_key,
&worker_id_label,
command_name,
"error",
);
tracing::warn!(
worker_id,
%error,
"Job command failed and will be retried"
);
delivery.nack_requeue().await?;
}
}
}
Err(err) => {
#[cfg(feature = "prometheus")]
metrics::record_command(routing_key, &worker_id_label, "invalid", "discarded");
tracing::warn!(
"[Worker#{}] Unknown message received [{} bytes]: {}",
worker_id,
delivery.data().len(),
err
);
delivery.ack().await?;
}
},
Err(e) => {
tracing::warn!("[Worker#{}] Consumer ended: {:?}", worker_id, e);
}
}
}
}
impl<C, H> Drop for BackgroundJobServer<C, H>
where
C: Sync + Send + 'static,
H: BgJobHandler<C> + Sync + Send + 'static,
{
fn drop(&mut self) {
match self.worker_tasks.lock() {
Ok(workers) => {
for worker in workers.iter() {
worker.task.abort();
}
}
Err(poisoned) => {
for worker in poisoned.into_inner().iter() {
worker.task.abort();
}
}
}
for task in &self.maintenance_tasks {
task.abort();
}
if let Some(task) = &self.partition_membership_task {
task.abort();
}
match self.partition_tasks.lock() {
Ok(tasks) => {
for task in tasks.iter() {
task.abort();
}
}
Err(poisoned) => {
for task in poisoned.into_inner().iter() {
task.abort();
}
}
}
if let Some(owner) = self.partition_owner.take() {
if let Ok(handle) = tokio::runtime::Handle::try_current() {
let handler = self.handler.clone();
handle.spawn(async move {
if let Err(error) = handler
.get_publisher()
.deregister_partition_worker(&owner)
.await
{
tracing::warn!(%error, "Failed to deregister partition worker on shutdown");
}
});
}
}
}
}