use std::{
pin::Pin,
sync::{
Arc,
atomic::{AtomicBool, Ordering},
},
task::{Context, Poll},
time::Duration,
};
use diesel::{RunQueryDsl, sql_query};
use futures::{Stream, channel::mpsc};
use crate::{Error, InsertEvent, PgPool, PgTaskId};
pub(crate) const NOTIFY_LISTENER_POLL_INTERVAL: Duration = Duration::from_millis(50);
pub(crate) const NOTIFY_CHANNEL_CAPACITY_MAX: usize = 8192;
pub(crate) fn clamp_notify_capacity(capacity: usize) -> usize {
capacity.clamp(1, NOTIFY_CHANNEL_CAPACITY_MAX)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum DeliveryOutcome {
Delivered,
ChannelFull,
ReceiverGone,
}
fn classify_delivery<T>(result: Result<(), mpsc::TrySendError<T>>) -> DeliveryOutcome {
match result {
Ok(()) => DeliveryOutcome::Delivered,
Err(error) if error.is_disconnected() => DeliveryOutcome::ReceiverGone,
Err(_) => DeliveryOutcome::ChannelFull,
}
}
pub(crate) fn notify_task_ids(
pool: PgPool,
queue: String,
capacity: usize,
) -> impl Stream<Item = Result<PgTaskId, Error>> + Send {
let (mut sender, receiver) = mpsc::channel(clamp_notify_capacity(capacity));
let cancel = Arc::new(AtomicBool::new(false));
let thread_cancel = cancel.clone();
let mut spawn_error_sender = sender.clone();
let thread_pool = pool.clone();
if let Err(error) = std::thread::Builder::new()
.name("apalis-postgres-notify".to_owned())
.spawn(move || {
let mut conn = match thread_pool.get() {
Ok(conn) => conn,
Err(error) => {
let _ = sender.try_send(Err(Error::from(error)));
return;
}
};
if let Err(error) = sql_query("LISTEN \"apalis::job::insert\"").execute(&mut conn) {
let _ = sender.try_send(Err(Error::database(
"starting PostgreSQL LISTEN notification listener",
)(error)));
return;
}
let unlisten = |conn: &mut diesel::PgConnection| {
let _ = sql_query("UNLISTEN \"apalis::job::insert\"").execute(conn);
};
'listen: while !thread_cancel.load(Ordering::Acquire) {
for notification in conn.notifications_iter() {
if thread_cancel.load(Ordering::Acquire) {
break 'listen;
}
let notification = match notification {
Ok(notification) => notification,
Err(error) => {
let _ = sender.try_send(Err(Error::database(
"receiving PostgreSQL notification",
)(error)));
break 'listen;
}
};
let Ok(event) = serde_json::from_str::<InsertEvent>(¬ification.payload)
else {
continue;
};
let (event_queue, ids) = event.into_ids();
if event_queue != queue {
continue;
}
for id in ids {
match classify_delivery(sender.try_send(Ok(id))) {
DeliveryOutcome::Delivered => {}
DeliveryOutcome::ReceiverGone => break 'listen,
DeliveryOutcome::ChannelFull => break,
}
}
}
std::thread::sleep(NOTIFY_LISTENER_POLL_INTERVAL);
}
unlisten(&mut conn);
})
{
let _ = spawn_error_sender.try_send(Err(Error::NotifyListener(error.to_string())));
}
NotifyTaskIds {
receiver,
cancel,
pool,
}
}
pub(crate) struct NotifyTaskIds {
receiver: mpsc::Receiver<Result<PgTaskId, Error>>,
cancel: Arc<AtomicBool>,
pool: PgPool,
}
impl Stream for NotifyTaskIds {
type Item = Result<PgTaskId, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Pin::new(&mut self.receiver).poll_next(cx)
}
}
impl Drop for NotifyTaskIds {
fn drop(&mut self) {
self.cancel.store(true, Ordering::Release);
let pool = self.pool.clone();
let _ = std::thread::Builder::new()
.name("apalis-postgres-notify-drop".to_owned())
.spawn(move || {
if let Ok(mut conn) = pool.get() {
let _ =
sql_query("SELECT pg_notify('apalis::job::insert', '')").execute(&mut conn);
}
});
}
}
#[cfg(test)]
mod tests {
use lets_expect::*;
use super::*;
fn classify_a_delivered_send() -> DeliveryOutcome {
let (mut sender, _receiver) = mpsc::channel::<i32>(1);
classify_delivery(sender.try_send(1))
}
fn classify_a_full_channel() -> DeliveryOutcome {
let (mut sender, _receiver) = mpsc::channel::<i32>(1);
while sender.try_send(1).is_ok() {}
classify_delivery(sender.try_send(1))
}
fn classify_a_dropped_receiver() -> DeliveryOutcome {
let (mut sender, receiver) = mpsc::channel::<i32>(1);
drop(receiver);
classify_delivery(sender.try_send(1))
}
fn clamp_capacity(capacity: usize) -> usize {
clamp_notify_capacity(capacity)
}
lets_expect! {
expect(clamp_capacity(capacity)) {
let capacity = 8_usize;
to preserves_the_caller_value { equal(8) }
when the_capacity_is_below_the_minimum {
let capacity = 0_usize;
to clamps_up_to_the_minimum_of_one { equal(1) }
}
when the_capacity_equals_the_maximum {
let capacity = NOTIFY_CHANNEL_CAPACITY_MAX;
to keeps_the_maximum_unchanged { equal(NOTIFY_CHANNEL_CAPACITY_MAX) }
}
when the_capacity_exceeds_the_maximum {
let capacity = NOTIFY_CHANNEL_CAPACITY_MAX + 1;
to clamps_down_to_the_channel_capacity_cap {
equal(NOTIFY_CHANNEL_CAPACITY_MAX)
}
}
}
expect(classify_a_delivered_send()) {
when the_channel_accepts_the_id {
to reports_the_id_as_delivered { equal(DeliveryOutcome::Delivered) }
}
}
expect(classify_a_full_channel()) {
when the_channel_is_full_but_still_connected {
to drops_the_wakeup_without_stopping_the_listener {
equal(DeliveryOutcome::ChannelFull)
}
}
}
expect(classify_a_dropped_receiver()) {
when the_receiver_has_been_dropped {
to signals_the_listener_to_stop { equal(DeliveryOutcome::ReceiverGone) }
}
}
}
}