feat(transport): Worker eviction (#185)

Adds max worker capacity of 500 to the worker pool,
and an automatic shutdown of idle workers.
Messages that would cause the capacity to be exceeded,
are deferred.

Additionally, ensures that the same worker
is not spawned by two tasks at the same time.

Signed-off-by: Jagoda Ślązak <jslazak@jslazak.com>
This commit is contained in:
Jagoda Estera Ślązak
2026-06-27 10:54:38 +02:00
committed by GitHub
parent ac9df36801
commit d5ac67ecff
5 changed files with 209 additions and 47 deletions
+1 -1
View File
@@ -4,4 +4,4 @@ pub const OK_SMTP_250: &str = "250 OK (SMTP)";
pub const ENCRYPTION_NEEDED_523: &str = "523 Encryption Needed: Invalid Unencrypted Mail"; pub const ENCRYPTION_NEEDED_523: &str = "523 Encryption Needed: Invalid Unencrypted Mail";
pub const LOCAL_ERROR_451: &str = "451 Local error"; pub const LOCAL_ERROR_451: &str = "451 Local error";
pub const WORKER_BUSY_421: &str = "421 Worker for this destination is busy"; pub const TRANSPORT_BUSY_421: &str = "421 Transport busy";
@@ -20,7 +20,7 @@ TRANSACTION 1
> message > message
> . > .
< 250 OK (SMTP) < 250 OK (SMTP)
421 Worker for this destination is busy 421 Transport busy
[filtermail-transport -> destination A] [filtermail-transport -> destination A]
< 220 filtermail SMTP < 220 filtermail SMTP
@@ -0,0 +1,51 @@
---
source: src/transport.rs
expression: "format!(\"TRANSACTION 1\\r\\n\\\n [postfix -> filtermail-transport]\\r\\n{record_postfix_1}\\r\\n\\\n [filtermail-transport -> destination A]\\r\\n{record_filtermail_1}\\r\\n\\r\\n\\\n TRANSACTION 2\\r\\n\\\n [postfix -> filtermail-transport]\\r\\n{record_postfix_2}\")"
---
TRANSACTION 1
[postfix -> filtermail-transport]
< 220 filtermail SMTP
> LHLO postfix
< 250-filtermail
250-8BITMIME
250 OK
> MAIL FROM:<sender@here>
< 250 OK
> RCPT TO:<a1@localhost>
< 250 OK
> DATA
< 354 End data with <CR><LF>.<CR><LF>
> message
> .
< 250 OK (SMTP)
[filtermail-transport -> destination A]
< 220 filtermail SMTP
> EHLO example.org
< 250-filtermail
250-8BITMIME
250 OK
> MAIL FROM:<sender@here>
< 250 OK
> RCPT TO:<a1@localhost>
< 250 OK
> DATA
< 354 End data with <CR><LF>.<CR><LF>
> message
.
< 250 OK
TRANSACTION 2
[postfix -> filtermail-transport]
< 220 filtermail SMTP
> LHLO postfix
< 250-filtermail
250-8BITMIME
250 OK
> MAIL FROM:<sender@here>
< 250 OK
> RCPT TO:<b1@[127.0.0.1]>
< 250 OK
> DATA
< 421 Transport busy
+76 -15
View File
@@ -2,7 +2,7 @@ mod https_client;
mod worker; mod worker;
use crate::config::Config; use crate::config::Config;
use crate::smtp_responses::{LOCAL_ERROR_451, WORKER_BUSY_421}; use crate::smtp_responses::{LOCAL_ERROR_451, TRANSPORT_BUSY_421};
use crate::smtp_server::{SmtpHandler, Transaction}; use crate::smtp_server::{SmtpHandler, Transaction};
use crate::tcp::{TcpConnect, TcpStreamTrait}; use crate::tcp::{TcpConnect, TcpStreamTrait};
use crate::utils::AddressDomain; use crate::utils::AddressDomain;
@@ -32,12 +32,17 @@ where
Ok(Self { workers }) Ok(Self { workers })
} }
/// Same as [`Self::new`], but lets you set worker queue size. /// Same as [`Self::new`], but lets you set worker queue size and pool capacity.
/// ///
/// Only used for tests. /// Only used for tests.
#[cfg(test)] #[cfg(test)]
pub fn with_queue_size(config: Config, queue_size: usize) -> Result<Self, crate::error::Error> { pub fn with_queue_size_and_pool_capacity(
let workers = WorkerPool::with_queue_size(config, queue_size)?; config: Config,
queue_size: usize,
pool_capacity: usize,
) -> Result<Self, crate::error::Error> {
let workers =
WorkerPool::with_queue_size_and_pool_capacity(config, queue_size, pool_capacity)?;
Ok(Self { workers }) Ok(Self { workers })
} }
@@ -94,7 +99,7 @@ where
// deferred mails to be constantly retried. // deferred mails to be constantly retried.
if transaction.state.permits.is_empty() { if transaction.state.permits.is_empty() {
return Err(WORKER_BUSY_421.to_string()); return Err(TRANSPORT_BUSY_421.to_string());
} }
Ok(()) Ok(())
@@ -133,11 +138,10 @@ where
if let Some(permit) = transaction.state.permits.remove(rcpt_domain) { if let Some(permit) = transaction.state.permits.remove(rcpt_domain) {
let (message, receiver) = WorkerMessage::new(domain_envelope); let (message, receiver) = WorkerMessage::new(domain_envelope);
permit.send(message); permit.send(message);
// todo: receiver timeout?
transactions.spawn(receiver).id() transactions.spawn(receiver).id()
} else { } else {
transactions transactions
.spawn(async move { Ok(Err(WORKER_BUSY_421.to_string())) }) .spawn(async move { Ok(Err(TRANSPORT_BUSY_421.to_string())) })
.id() .id()
}; };
task_id_domain_map.insert(receiver_task_id, rcpt_domain); task_id_domain_map.insert(receiver_task_id, rcpt_domain);
@@ -152,11 +156,14 @@ where
let smtp_response = match result { let smtp_response = match result {
Ok((_, Ok(Ok(resp)))) | Ok((_, Ok(Err(resp)))) => resp, Ok((_, Ok(Ok(resp)))) | Ok((_, Ok(Err(resp)))) => resp,
Ok((_, Err(e))) => { Ok((_, Err(_))) => {
log::error!( log::error!(
"Worker task failed while delivering to {}: {e}", "Delivery to {} failed due to a dead worker!",
domain.map(AsRef::as_ref).unwrap_or("<unknown>") domain.map(AsRef::as_ref).unwrap_or("<unknown>")
); );
if let Some(domain) = domain {
self.workers.cleanup(domain);
}
LOCAL_ERROR_451.to_string() LOCAL_ERROR_451.to_string()
} }
Err(e) => { Err(e) => {
@@ -243,9 +250,16 @@ mod tests {
/// Spawns filtermail-transport. /// Spawns filtermail-transport.
/// ///
/// Returns a pointer to the underlying handler. /// Returns a pointer to the underlying handler.
fn spawn_filtermail_transport() -> TestResult<Arc<TransportHandler<TcpStream>>> { fn spawn_filtermail_transport(
queue_size: usize,
pool_capacity: usize,
) -> TestResult<Arc<TransportHandler<TcpStream>>> {
let config = Config::default(); let config = Config::default();
let transport = Arc::new(TransportHandler::with_queue_size(config.clone(), 1)?); let transport = Arc::new(TransportHandler::with_queue_size_and_pool_capacity(
config.clone(),
queue_size,
pool_capacity,
)?);
tokio::spawn(run_smtp_server( tokio::spawn(run_smtp_server(
&FILTERMAIL_ADDR, &FILTERMAIL_ADDR,
transport.clone(), transport.clone(),
@@ -302,8 +316,11 @@ mod tests {
#[rstest] #[rstest]
#[tokio::test] #[tokio::test]
async fn test_rcpt_to_and_start_data(addrs1: Vec<String>, addrs2: Vec<String>) -> TestResult { async fn test_rcpt_to_and_start_data(addrs1: Vec<String>, addrs2: Vec<String>) -> TestResult {
let transport_handler = let transport_handler = TransportHandler::<TcpStream>::with_queue_size_and_pool_capacity(
TransportHandler::<TcpStream>::with_queue_size(Config::default(), 1)?; Config::default(),
1,
500,
)?;
let domain1 = AddressDomain::from_str(addrs1.first().unwrap())?; let domain1 = AddressDomain::from_str(addrs1.first().unwrap())?;
let domain2 = AddressDomain::from_str(addrs2.first().unwrap())?; let domain2 = AddressDomain::from_str(addrs2.first().unwrap())?;
@@ -358,7 +375,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_smtp_send_mail() -> TestResult { async fn test_smtp_send_mail() -> TestResult {
let mut remote_mta = spawn_mock_mta()?; let mut remote_mta = spawn_mock_mta()?;
spawn_filtermail_transport()?; spawn_filtermail_transport(1, 2)?;
let envelope = Envelope { let envelope = Envelope {
mail_from: "sender@here".to_string(), mail_from: "sender@here".to_string(),
@@ -398,7 +415,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_smtp_send_mail_defer() -> TestResult { async fn test_smtp_send_mail_defer() -> TestResult {
let mut remote_mta = spawn_mock_mta()?; let mut remote_mta = spawn_mock_mta()?;
let transport = spawn_filtermail_transport()?; let transport = spawn_filtermail_transport(1, 500)?;
let mut envelope = Envelope { let mut envelope = Envelope {
mail_from: "sender@here".to_string(), mail_from: "sender@here".to_string(),
@@ -438,4 +455,48 @@ mod tests {
Ok(()) Ok(())
} }
#[rstest]
#[serial]
#[tokio::test]
async fn test_smtp_send_mail_worker_cap_exceeded() -> TestResult {
let mut remote_mta = spawn_mock_mta()?;
spawn_filtermail_transport(30, 1)?;
let envelope_1 = Envelope {
mail_from: "sender@here".to_string(),
rcpt_to: vec!["a1@localhost".to_string()],
data: "message\r\n".as_bytes().to_vec(),
};
let envelope_2 = Envelope {
mail_from: "sender@here".to_string(),
rcpt_to: vec!["b1@[127.0.0.1]".to_string()],
data: "message\r\n".as_bytes().to_vec(),
};
// 1
let (record_postfix_1, record_filtermail_1) = {
let record_postfix = lmtp_send(&envelope_1).await?;
let record_filtermail = remote_mta.recv().await.unwrap();
tokio::time::sleep(Duration::from_secs(1)).await;
assert!(remote_mta.is_empty());
(record_postfix, record_filtermail)
};
// 2
let record_postfix_2 = lmtp_send(&envelope_2).await?;
tokio::time::sleep(Duration::from_secs(1)).await;
assert!(remote_mta.is_empty());
insta::assert_snapshot!(format!(
"TRANSACTION 1\r\n\
[postfix -> filtermail-transport]\r\n{record_postfix_1}\r\n\
[filtermail-transport -> destination A]\r\n{record_filtermail_1}\r\n\r\n\
TRANSACTION 2\r\n\
[postfix -> filtermail-transport]\r\n{record_postfix_2}"
));
Ok(())
}
} }
+80 -30
View File
@@ -17,6 +17,7 @@ use tokio::sync::mpsc::OwnedPermit;
use tokio::sync::{mpsc, oneshot}; use tokio::sync::{mpsc, oneshot};
use tokio::task; use tokio::task;
use tokio::task::JoinHandle; use tokio::task::JoinHandle;
use tokio::time::timeout;
use tokio_rustls::rustls; use tokio_rustls::rustls;
#[cfg(not(test))] #[cfg(not(test))]
@@ -35,16 +36,25 @@ const SMTP_SKIP_TLS: bool = true;
/// all new messages will be immediately deferred. /// all new messages will be immediately deferred.
const PER_DESTINATION_QUEUE_SIZE: usize = 30; const PER_DESTINATION_QUEUE_SIZE: usize = 30;
/// Max number of [`Worker`]s operating at the same time.
const MAX_WORKERS: usize = 500;
/// How long a worker can stay idle.
///
/// Exceeding this value will cause a worker shutdown.
const WORKER_KEEPALIVE_DURATION: Duration = Duration::from_secs(60);
type SMTPResponse = Result<String, String>; type SMTPResponse = Result<String, String>;
pub struct WorkerPool<S: TcpConnect> { pub struct WorkerPool<S: TcpConnect> {
inner: RwLock<BTreeMap<AddressDomain, Arc<Worker>>>, inner: Arc<RwLock<BTreeMap<AddressDomain, Arc<Worker>>>>,
client_hostname: String, client_hostname: String,
smtp_connection_pool: Arc<SmtpConnectionPool<S>>, smtp_connection_pool: Arc<SmtpConnectionPool<S>>,
mxdeliv_unsupported_hosts: Arc<retainer::Cache<String, ()>>, mxdeliv_unsupported_hosts: Arc<retainer::Cache<String, ()>>,
monitor_handle: JoinHandle<()>, monitor_handle: JoinHandle<()>,
dns_resolver: Arc<TokioResolver>, dns_resolver: Arc<TokioResolver>,
queue_size: usize, queue_size: usize,
capacity: usize,
} }
impl<S> WorkerPool<S> impl<S> WorkerPool<S>
@@ -72,48 +82,61 @@ where
mxdeliv_unsupported_hosts: mxdeliv_cache, mxdeliv_unsupported_hosts: mxdeliv_cache,
monitor_handle, monitor_handle,
queue_size: PER_DESTINATION_QUEUE_SIZE, queue_size: PER_DESTINATION_QUEUE_SIZE,
capacity: MAX_WORKERS,
}) })
} }
/// Same as [`Self::new`], but lets you set the size of the queue. /// Same as [`Self::new`], but lets you set the size of the queue and pool capacity.
/// ///
/// Used only for tests. /// Used only for tests.
#[cfg(test)] #[cfg(test)]
pub fn with_queue_size(config: Config, queue_size: usize) -> Result<Self, crate::error::Error> { pub fn with_queue_size_and_pool_capacity(
config: Config,
queue_size: usize,
pool_capacity: usize,
) -> Result<Self, crate::error::Error> {
let mut this = Self::new(config)?; let mut this = Self::new(config)?;
this.queue_size = queue_size; this.queue_size = queue_size;
this.capacity = pool_capacity;
Ok(this) Ok(this)
} }
fn get_or_create_worker(&self, destination: &AddressDomain) -> Arc<Worker> { /// Gets a worker for provided `destination`,
/// spawning a new one if required.
///
/// Returns [`None`], if operation requires spawning a new worker,
/// but pool already operates at maximum worker capacity.
fn get_or_create_worker(&self, destination: &AddressDomain) -> Option<Arc<Worker>> {
// NOTE: these locks are blocking, but critical section here is quite small and // NOTE: these locks are blocking, but critical section here is quite small and
// shouldn't cause issues in async code. // shouldn't cause issues in async code.
// NOTE: read() returns a guard that is dropped before the match statement. // NOTE: read() returns a guard that is dropped before the match statement.
// This must be ensured or else, the write() line would cause a deadlock. // This must be ensured or else, the write() line would cause a deadlock.
let worker = { let worker = {
let mut worker = { let map = self.inner.read();
let map = self.inner.read(); map.get(destination).cloned()
map.get(destination).cloned()
};
// Remove (and re-create) worker if it finished/crashed.
// In reality, this should never happen.
if let Some(w) = &worker
&& w.handle.is_finished()
{
log::error!("Worker for destination {destination} crashed! Restarting...",);
worker = None;
{
let mut map = self.inner.write();
map.remove(destination);
}
};
worker
}; };
match worker { match worker {
Some(worker) => worker, Some(worker) => Some(worker),
None => { None => {
// Worker for this destination wasn't spawned yet. // Worker for this destination wasn't spawned yet.
let mut map = self.inner.write();
// we check if it wasn't spawned in the meantime first
if let Some(worker) = map.get(destination)
&& !worker.handle.is_finished()
{
return Some(worker.clone());
}
if map.len() >= self.capacity {
log::warn!(
"Worker pool operating at maximum capacity! \
Messages to new destinations will be deferred."
);
return None;
}
let (tx, rx) = mpsc::channel(self.queue_size); let (tx, rx) = mpsc::channel(self.queue_size);
let handle = tokio::spawn(Worker::run( let handle = tokio::spawn(Worker::run(
destination.clone(), destination.clone(),
@@ -122,24 +145,47 @@ where
self.smtp_connection_pool.clone(), self.smtp_connection_pool.clone(),
self.mxdeliv_unsupported_hosts.clone(), self.mxdeliv_unsupported_hosts.clone(),
self.dns_resolver.clone(), self.dns_resolver.clone(),
self.inner.clone(),
)); ));
log::trace!("Worker {} spawned", handle.id()); log::trace!("Worker {} spawned", handle.id());
let worker = Arc::new(Worker { tx, handle }); let worker = Arc::new(Worker { tx, handle });
self.inner map.insert(destination.clone(), worker.clone());
.write() Some(worker)
.insert(destination.clone(), worker.clone());
worker
} }
} }
} }
/// Tries to get an [`OwnedPermit`] to the worker for specified destination. /// Tries to get an [`OwnedPermit`] to the worker for specified destination.
/// ///
/// Returns [`None`] if the worker's queue is full. /// Returns [`None`] if:
///
/// - the worker's queue is full,
/// - operation requires spawning a new worker,
/// but pool already operates at maximum worker capacity.
pub fn get_permit(&self, destination: &AddressDomain) -> Option<OwnedPermit<WorkerMessage>> { pub fn get_permit(&self, destination: &AddressDomain) -> Option<OwnedPermit<WorkerMessage>> {
let worker = self.get_or_create_worker(destination); if let Some(worker) = self.get_or_create_worker(destination) {
worker.tx.clone().try_reserve_owned().ok() return worker.tx.clone().try_reserve_owned().ok();
}
None
}
/// Informs pool that a worker has exited unexpectedly.
///
/// Such worker will be removed from the pool,
/// and re-created on the next [`Self::get_permit`] call.
pub fn cleanup(&self, destination: &AddressDomain) {
let mut map = self.inner.write();
// we first check in case other task already removed/re-created it
if let Some(worker) = map.get(destination)
&& worker.handle.is_finished()
{
log::info!(
"Removing a dead worker {} for destination {destination}",
worker.handle.id()
);
map.remove(destination);
}
} }
} }
@@ -169,6 +215,7 @@ impl Worker {
smtp_connection_pool: Arc<SmtpConnectionPool<S>>, smtp_connection_pool: Arc<SmtpConnectionPool<S>>,
mxdeliv_unsupported_hosts: Arc<retainer::Cache<String, ()>>, mxdeliv_unsupported_hosts: Arc<retainer::Cache<String, ()>>,
dns_resolver: Arc<TokioResolver>, dns_resolver: Arc<TokioResolver>,
worker_pool: Arc<RwLock<BTreeMap<AddressDomain, Arc<Worker>>>>,
) -> Result<(), crate::error::Error> ) -> Result<(), crate::error::Error>
where where
S: TcpStreamTrait + TcpConnect, S: TcpStreamTrait + TcpConnect,
@@ -182,7 +229,7 @@ impl Worker {
let tls_resumption_store = Arc::new(rustls::client::ClientSessionMemoryCache::new(256)); let tls_resumption_store = Arc::new(rustls::client::ClientSessionMemoryCache::new(256));
let https_client = HttpsClient::new(tls_resumption_store.clone())?; let https_client = HttpsClient::new(tls_resumption_store.clone())?;
while let Some(message) = rx.recv().await { while let Ok(Some(message)) = timeout(WORKER_KEEPALIVE_DURATION, rx.recv()).await {
log::trace!( log::trace!(
"Worker {worker_id} received a message from {}", "Worker {worker_id} received a message from {}",
message.envelope.mail_from message.envelope.mail_from
@@ -205,6 +252,9 @@ impl Worker {
}; };
} }
log::info!("Worker {worker_id} for domain {destination} shutting down...");
worker_pool.write().remove(&destination);
Ok(()) Ok(())
} }