diff --git a/crates/kumod/src/http_server/inject_v1.rs b/crates/kumod/src/http_server/inject_v1.rs index 95a601ab..4b7f9d23 100644 --- a/crates/kumod/src/http_server/inject_v1.rs +++ b/crates/kumod/src/http_server/inject_v1.rs @@ -561,7 +561,7 @@ async fn process_recipient<'a>( request.trace_headers.apply_supplemental(&message)?; if !request.deferred_spool { - message.save().await?; + message.save(None).await?; } log_disposition(LogDisposition { kind: RecordType::Reception, @@ -617,7 +617,7 @@ async fn queue_deferred( message.set_meta("received_from", peer_address.to_string())?; message.set_meta("queue", GENERATOR_QUEUE_NAME)?; if !request.deferred_spool { - message.save().await?; + message.save(None).await?; } log_disposition(LogDisposition { kind: RecordType::Reception, diff --git a/crates/kumod/src/logging/hooks.rs b/crates/kumod/src/logging/hooks.rs index 890f4fe6..87ee6e67 100644 --- a/crates/kumod/src/logging/hooks.rs +++ b/crates/kumod/src/logging/hooks.rs @@ -175,7 +175,7 @@ impl LogHookState { if enqueue { let queue_name = msg.get_queue_name()?; if !deferred_spool { - msg.save().await?; + msg.save(None).await?; } QueueManager::insert(&queue_name, msg, InsertReason::Received.into()).await?; } diff --git a/crates/kumod/src/queue.rs b/crates/kumod/src/queue.rs index 7edf6d6b..d16d862b 100644 --- a/crates/kumod/src/queue.rs +++ b/crates/kumod/src/queue.rs @@ -45,6 +45,7 @@ use thiserror::Error; use throttle::{ThrottleResult, ThrottleSpec}; use timeq::{PopResult, TimeQ, TimerEntryWithDelay, TimerError}; use tokio::sync::{Notify, Semaphore}; +use tokio::time::timeout_at; use tracing::instrument; static MANAGER: LazyLock = LazyLock::new(|| QueueManager::new()); @@ -1391,7 +1392,7 @@ impl Queue { } if msg.needs_save() { - if let Err(err) = msg.save().await { + if let Err(err) = msg.save(None).await { tracing::error!("failed to save msg after rebind: {err:#}"); } } @@ -1707,7 +1708,7 @@ impl Queue { ) -> anyhow::Result<()> { tracing::trace!("save_if_needed {}", msg.id()); if msg.needs_save() { - msg.save().await?; + msg.save(None).await?; } match queue_config { @@ -2725,15 +2726,19 @@ impl QueueManager { deadline: Option, ) -> anyhow::Result<()> { tracing::trace!("QueueManager::insert {context:?}"); - let timer = RESOLVE_LATENCY.start_timer(); - let entry = Self::resolve(name).await?; - timer.stop_and_record(); - if let Some(deadline) = deadline { - if deadline <= Instant::now() { - anyhow::bail!("data processing deadline exceeded, stopping short of insertion"); + let timer = RESOLVE_LATENCY.start_timer(); + let entry = if let Some(deadline) = deadline { + match timeout_at(deadline.into(), Self::resolve(name)).await { + Err(_) => { + anyhow::bail!("data processing deadline exceeded, stopping short of insertion") + } + Ok(result) => result?, } - } + } else { + Self::resolve(name).await? + }; + timer.stop_and_record(); let _timer = INSERT_LATENCY.start_timer(); entry.insert(msg, context).await diff --git a/crates/kumod/src/smtp_server.rs b/crates/kumod/src/smtp_server.rs index 70c0a84b..f5878ea5 100644 --- a/crates/kumod/src/smtp_server.rs +++ b/crates/kumod/src/smtp_server.rs @@ -1741,7 +1741,7 @@ impl SmtpServer { Err(_) => { self.write_response( 451, - "4.4.5 data_processing_timeout exceeded", + "4.4.5 data_processing_timeout exceeded (rx)", Some("DATA".into()), ) .await?; @@ -1773,6 +1773,29 @@ impl SmtpServer { let mut was_arf_or_oob = false; let mut black_holed = false; + // pre-resolve any queues; there can be DNS and other async components + // to resolution that can cause this to take a non-trivial amount of time, + // so let's get that out of the way before we start writing to spool, + // to make it less complex to unwind if we exceed the allowed time. + for message in &accepted_messages { + let queue_name = message.get_queue_name()?; + match timeout_at(deadline.into(), QueueManager::resolve(&queue_name)).await { + Err(_) => { + self.write_response( + 451, + "4.4.5 data_processing_timeout exceeded (resolve)", + Some("DATA".into()), + ) + .await?; + return Ok(()); + } + Ok(Err(err)) => { + anyhow::bail!("QueueManager::resolve({queue_name}) failed: {err:#}"); + } + Ok(Ok(_handle)) => {} + } + } + for message in accepted_messages { self.params.trace_headers.apply_supplemental(&message)?; @@ -1807,7 +1830,23 @@ impl SmtpServer { if queue_name != "null" { if relay_disposition.relay && !self.params.deferred_spool { - message.save().await?; + match message.save(Some(deadline)).await { + Err(err) => { + // FIXME: unwind rest of batch + + if err.root_cause().is::() { + self.write_response( + 451, + "4.4.5 data_processing_timeout exceeded (spool)", + Some("DATA".into()), + ) + .await?; + return Ok(()); + } + return Err(err); + } + Ok(_) => {} + } } } diff --git a/crates/message/src/message.rs b/crates/message/src/message.rs index 2aee8e30..be663838 100644 --- a/crates/message/src/message.rs +++ b/crates/message/src/message.rs @@ -25,7 +25,7 @@ use serde::{Deserialize, Serialize}; use spool::{get_data_spool, get_meta_spool, Spool, SpoolId}; use std::hash::Hash; use std::sync::{Arc, LazyLock, Mutex, Weak}; -use std::time::Duration; +use std::time::{Duration, Instant}; use timeq::TimerEntryWithDelay; bitflags::bitflags! { @@ -477,15 +477,17 @@ impl Message { || inner.flags.contains(MessageFlags::DATA_DIRTY) } - pub async fn save(&self) -> anyhow::Result<()> { + pub async fn save(&self, deadline: Option) -> anyhow::Result<()> { let _timer = SAVE_HIST.start_timer(); - self.save_to(&**get_meta_spool(), &**get_data_spool()).await + self.save_to(&**get_meta_spool(), &**get_data_spool(), deadline) + .await } pub async fn save_to( &self, meta_spool: &(dyn Spool + Send + Sync), data_spool: &(dyn Spool + Send + Sync), + deadline: Option, ) -> anyhow::Result<()> { let force_sync = self .msg_and_id @@ -498,7 +500,7 @@ impl Message { let data_fut = if let Some(data) = self.get_data_if_dirty() { anyhow::ensure!(!data.is_empty(), "message data must not be empty"); data_spool - .store(self.msg_and_id.id, data, force_sync) + .store(self.msg_and_id.id, data, force_sync, deadline) .map(|_| true) .boxed() } else { @@ -507,13 +509,18 @@ impl Message { let meta_fut = if let Some(meta) = self.get_meta_if_dirty() { let meta = Arc::new(serde_json::to_vec(&meta)?.into_boxed_slice()); meta_spool - .store(self.msg_and_id.id, meta, force_sync) + .store(self.msg_and_id.id, meta, force_sync, deadline) .map(|_| true) .boxed() } else { futures::future::ready(false).boxed() }; + // NOTE: if we have a deadline, it is tempting to want to use + // timeout_at here to enforce it, but the underlying spool + // futures are not guaranteed to be fully cancel safe, which + // is why we pass the deadline down to the save method to allow + // them to handle timeouts internally. let (data_res, meta_res) = tokio::join!(data_fut, meta_fut); if data_res { @@ -541,13 +548,13 @@ impl Message { /// Save the data+meta if needed, then release both pub async fn save_and_shrink(&self) -> anyhow::Result { - self.save().await?; + self.save(None).await?; self.shrink() } /// Save the data+meta if needed, then release just the data pub async fn save_and_shrink_data(&self) -> anyhow::Result { - self.save().await?; + self.save(None).await?; self.shrink_data() } @@ -1256,14 +1263,14 @@ impl UserData for Message { methods.add_async_method("shrink", |_, this, _: ()| async move { if this.needs_save() { - this.save().await.map_err(any_err)?; + this.save(None).await.map_err(any_err)?; } this.shrink().map_err(any_err) }); methods.add_async_method("shrink_data", |_, this, _: ()| async move { if this.needs_save() { - this.save().await.map_err(any_err)?; + this.save(None).await.map_err(any_err)?; } this.shrink_data().map_err(any_err) }); @@ -1392,7 +1399,7 @@ impl UserData for Message { }); methods.add_async_method("save", |_, this, ()| async move { - this.save().await.map_err(any_err) + this.save(None).await.map_err(any_err) }); methods.add_method("set_force_sync", move |_, this, force: bool| { diff --git a/crates/spool/src/lib.rs b/crates/spool/src/lib.rs index 2e4841f2..b65d58ad 100644 --- a/crates/spool/src/lib.rs +++ b/crates/spool/src/lib.rs @@ -2,6 +2,7 @@ use async_trait::async_trait; use chrono::{DateTime, Utc}; use flume::Sender; use std::sync::{Arc, OnceLock}; +use std::time::Instant; pub mod local_disk; #[cfg(feature = "rocksdb")] @@ -30,6 +31,7 @@ pub trait Spool: Send + Sync { id: SpoolId, data: Arc>, force_sync: bool, + deadline: Option, ) -> anyhow::Result<()>; /// Scan the contents of the spool, and emit a SpoolEntry for each item diff --git a/crates/spool/src/local_disk.rs b/crates/spool/src/local_disk.rs index 1291ddd8..8dc4f75a 100644 --- a/crates/spool/src/local_disk.rs +++ b/crates/spool/src/local_disk.rs @@ -8,6 +8,7 @@ use std::io::Write; use std::os::fd::AsRawFd; use std::path::{Path, PathBuf}; use std::sync::Arc; +use std::time::Instant; use tempfile::NamedTempFile; use tokio::runtime::Handle; @@ -98,6 +99,7 @@ impl Spool for LocalDiskSpool { id: SpoolId, data: Arc>, force_sync: bool, + _deadline: Option, ) -> anyhow::Result<()> { let path = self.compute_path(id); let new_dir = self.path.join("new"); @@ -300,6 +302,7 @@ mod test { id, Arc::new(format!("I am {i}").as_bytes().to_vec().into_boxed_slice()), false, + None, ) .await?; ids.push(id); diff --git a/crates/spool/src/rocks.rs b/crates/spool/src/rocks.rs index 108728f6..234407aa 100644 --- a/crates/spool/src/rocks.rs +++ b/crates/spool/src/rocks.rs @@ -10,9 +10,10 @@ use rocksdb::{ use serde::{Deserialize, Serialize}; use std::path::Path; use std::sync::{Arc, LazyLock, Weak}; -use std::time::Duration; +use std::time::{Duration, Instant}; use tokio::runtime::Handle; use tokio::sync::Semaphore; +use tokio::time::timeout_at; #[derive(Serialize, Deserialize, Debug)] pub struct RocksSpoolParams { @@ -260,6 +261,7 @@ impl Spool for RocksSpool { id: SpoolId, data: Arc>, force_sync: bool, + deadline: Option, ) -> anyhow::Result<()> { let mut opts = WriteOptions::default(); opts.set_sync(force_sync); @@ -270,9 +272,12 @@ impl Spool for RocksSpool { match self.db.write_opt(batch, &opts) { Ok(()) => Ok(()), Err(err) if err.kind() == ErrorKind::Incomplete => { - let permit = match self.limit_concurrent_stores.clone() { - Some(s) => Some(s.acquire_owned().await?), - None => None, + let permit = match (self.limit_concurrent_stores.clone(), deadline) { + (Some(s), Some(deadline)) => { + Some(timeout_at(deadline.into(), s.acquire_owned()).await??) + } + (Some(s), None) => Some(s.acquire_owned().await?), + (None, _) => None, }; let db = self.db.clone(); tokio::task::Builder::new() @@ -441,6 +446,7 @@ mod test { id, Arc::new(format!("I am {i}").as_bytes().to_vec().into_boxed_slice()), false, + None, ) .await?; ids.push(id);