/* * SPDX-FileCopyrightText: 2020 Stalwart Labs LLC * * SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL */ use super::{ ArchivedMessage, ArchivedStatus, Message, MessageSource, QueueEnvelope, QueueId, QueuedMessage, QuotaKey, Recipient, Schedule, Status, }; use crate::queue::manager::{LockedMessage, Queue}; use crate::queue::{ FROM_AUTHENTICATED, FROM_AUTOGENERATED, FROM_DSN, FROM_REPORT, FROM_UNAUTHENTICATED, FROM_UNAUTHENTICATED_DMARC, MessageWrapper, }; use common::config::smtp::queue::QueueName; use common::ipc::QueueEvent; use common::{KV_LOCK_QUEUE_MESSAGE, Server}; use registry::pickle::Pickle; use registry::schema::prelude::{ObjectType, Property}; use registry::schema::structs::SpamTrainingSample; use registry::types::EnumImpl; use registry::types::datetime::UTCDateTime; use registry::types::id::ObjectId; use std::borrow::Cow; use std::collections::hash_map::Entry; use std::future::Future; use std::net::{IpAddr, Ipv4Addr}; use std::time::SystemTime; use store::write::key::DeserializeBigEndian; use store::write::serialize::rkyv_deserialize; use store::write::{ AlignedBytes, Archive, Archiver, BatchBuilder, BlobLink, BlobOp, MergeResult, Params, QueueClass, RegistryClass, ValueClass, now, }; use store::{ Deserialize, IterateParams, Serialize, SerializeInfallible, U64_LEN, ValueKey, xxhash_rust, }; use trc::{AddContext, ServerEvent, SpamEvent}; use types::blob::BlobId; use types::blob_hash::BlobHash; use utils::DomainPart; use utils::snowflake::SnowflakeIdGenerator; pub const LOCK_EXPIRY: u64 = 10 * 60; // 10 minutes pub const QUEUE_REFRESH: u64 = 5 * 60; // 5 minutes const INFINITE_LOCK: u64 = 60 * 60 * 24 * 365; // 1 year pub struct QueuedMessages { pub messages: Vec, pub next_refresh: u64, } pub trait SmtpSpool: Sync + Send { fn new_message(&self, return_path: impl AsRef, span_id: u64) -> MessageWrapper; fn next_event(&self, queue: &mut Queue) -> impl Future + Send; fn try_lock_event( &self, queue_id: QueueId, queue_name: QueueName, ) -> impl Future + Send; fn unlock_event( &self, queue_id: QueueId, queue_name: QueueName, ) -> impl Future + Send; fn read_message( &self, id: QueueId, queue_name: QueueName, ) -> impl Future> + Send; fn read_message_archive( &self, id: QueueId, ) -> impl Future>>> + Send; } impl SmtpSpool for Server { fn new_message(&self, return_path: impl AsRef, span_id: u64) -> MessageWrapper { let created = SystemTime::now() .duration_since(SystemTime::UNIX_EPOCH) .map_or(0, |d| d.as_secs()); MessageWrapper { queue_id: self.inner.data.queue_id_gen.generate(), queue_name: QueueName::default(), is_multi_queue: false, span_id, message: Message { created, return_path: return_path.to_lowercase_domain().into_boxed_str(), recipients: Vec::with_capacity(1), flags: 0, env_id: None, priority: 0, size: 0, blob_hash: Default::default(), quota_keys: Default::default(), received_from_ip: IpAddr::V4(Ipv4Addr::LOCALHOST), received_via_port: 0, }, } } async fn next_event(&self, queue: &mut Queue) -> QueuedMessages { let now = now(); let from_key = ValueKey::from(ValueClass::Queue(QueueClass::MessageEvent( store::write::QueueEvent { due: 0, queue_id: 0, queue_name: [0; 8], }, ))); let to_key = ValueKey::from(ValueClass::Queue(QueueClass::MessageEvent( store::write::QueueEvent { due: now + QUEUE_REFRESH, queue_id: u64::MAX, queue_name: [u8::MAX; 8], }, ))); let mut events = QueuedMessages { messages: Vec::new(), next_refresh: now + QUEUE_REFRESH, }; queue.locked_revision += 1; let result = self .store() .iterate( IterateParams::new(from_key, to_key).ascending().no_values(), |key, _| { let due = key.deserialize_be_u64(0)?; if due <= now { let queue_id = key.deserialize_be_u64(U64_LEN)?; let queue_name = key .get(U64_LEN + U64_LEN..) .and_then(QueueName::from_bytes) .ok_or_else(|| { trc::StoreEvent::DataCorruption .caused_by(trc::location!()) .ctx(trc::Key::Key, key) })?; let add_event = queue .stats .get(&queue_name) .is_none_or(|stats| stats.has_capacity()) && match queue.locked.entry((queue_id, queue_name)) { Entry::Occupied(mut entry) => { let locked = entry.get_mut(); locked.revision = queue.locked_revision; if locked.expires <= now { locked.expires = now + INFINITE_LOCK; true } else { if locked.expires < events.next_refresh { events.next_refresh = locked.expires; } false } } Entry::Vacant(entry) => { entry.insert(LockedMessage { expires: now + INFINITE_LOCK, revision: queue.locked_revision, }); true } }; if add_event { events.messages.push(QueuedMessage { due, queue_id, queue_name, }); } Ok(true) } else { if due < events.next_refresh { events.next_refresh = due; } Ok(false) } }, ) .await; if let Err(err) = result { trc::error!( err.details("Failed to read queue.") .caused_by(trc::location!()) ); } events } async fn try_lock_event(&self, queue_id: QueueId, queue_name: QueueName) -> bool { match self .in_memory_store() .try_lock( KV_LOCK_QUEUE_MESSAGE, &lock_id(queue_id, queue_name), LOCK_EXPIRY, ) .await { Ok(result) => { if !result { trc::event!( Queue(trc::QueueEvent::Locked), QueueId = queue_id, QueueName = queue_name.to_string() ); } result } Err(err) => { trc::error!( err.details("Failed to lock event.") .caused_by(trc::location!()) ); false } } } async fn unlock_event(&self, queue_id: QueueId, queue_name: QueueName) { if let Err(err) = self .in_memory_store() .remove_lock(KV_LOCK_QUEUE_MESSAGE, &lock_id(queue_id, queue_name)) .await { trc::error!( err.details("Failed to unlock event.") .caused_by(trc::location!()) ); } } async fn read_message( &self, queue_id: QueueId, queue_name: QueueName, ) -> Option { match self .read_message_archive(queue_id) .await .and_then(|a| match a { Some(a) => a.deserialize::().map(Some), None => Ok(None), }) { Ok(Some(message)) => Some(MessageWrapper { is_multi_queue: message.recipients.iter().any(|rcpt| { matches!(rcpt.status, Status::Scheduled | Status::TemporaryFailure(_)) && rcpt.queue != queue_name }), queue_id, queue_name, span_id: 0, message, }), Ok(None) => None, Err(err) => { trc::error!( err.details("Failed to read message.") .caused_by(trc::location!()) ); None } } } async fn read_message_archive( &self, id: QueueId, ) -> trc::Result>> { self.store() .get_value::>(ValueKey::from(ValueClass::Queue( QueueClass::Message(id), ))) .await } } fn lock_id(queue_id: QueueId, queue_name: QueueName) -> [u8; 16] { let mut id = [0; 16]; id[..8].copy_from_slice(&queue_id.to_be_bytes()); id[8..].copy_from_slice(queue_name.as_ref()); id } impl MessageWrapper { pub async fn queue( mut self, raw_headers: Option<&[u8]>, raw_message: &[u8], session_id: u64, server: &Server, source: MessageSource, ) -> bool { // Set flags let (flags, event, train_spam) = match source { MessageSource::Authenticated => ( FROM_AUTHENTICATED, trc::QueueEvent::QueueMessageAuthenticated, None, ), MessageSource::Unauthenticated { dmarc_pass: true, train_spam, } => ( FROM_UNAUTHENTICATED_DMARC, trc::QueueEvent::QueueMessage, train_spam, ), MessageSource::Unauthenticated { dmarc_pass: false, train_spam, } => ( FROM_UNAUTHENTICATED, trc::QueueEvent::QueueMessage, train_spam, ), MessageSource::Dsn => (FROM_DSN, trc::QueueEvent::QueueDsn, None), MessageSource::Report => (FROM_REPORT, trc::QueueEvent::QueueReport, None), MessageSource::Autogenerated => ( FROM_AUTOGENERATED, trc::QueueEvent::QueueAutogenerated, None, ), }; self.message.flags |= flags; // Write blob let message = if let Some(raw_headers) = raw_headers { let mut message = Vec::with_capacity(raw_headers.len() + raw_message.len()); message.extend_from_slice(raw_headers); message.extend_from_slice(raw_message); Cow::Owned(message) } else { raw_message.into() }; self.message.blob_hash = BlobHash::generate(message.as_ref()); // Generate id if self.message.size == 0 { self.message.size = message.len() as u64; } // Reserve and write blob let mut batch = BatchBuilder::new(); let now = now(); let reserve_until = now + 120; batch.set( BlobOp::Link { hash: self.message.blob_hash.clone(), to: BlobLink::Temporary { until: reserve_until, }, }, vec![], ); if let Err(err) = server.store().write(batch.build_all()).await { trc::error!( err.details("Failed to write to store.") .span_id(session_id) .caused_by(trc::location!()) ); return false; } if let Err(err) = server .blob_store() .put_blob( self.message.blob_hash.as_slice(), message.as_ref(), server.core.email.compression, ) .await { trc::error!( err.details("Failed to write blob.") .span_id(session_id) .caused_by(trc::location!()) ); return false; } trc::event!( Queue(event), SpanId = session_id, QueueId = self.queue_id, From = if !self.message.return_path.is_empty() { trc::Value::String(self.message.return_path.as_ref().into()) } else { trc::Value::String("<>".into()) }, To = self .message .recipients .iter() .map(|r| trc::Value::String(r.address.as_ref().into())) .collect::>(), Size = self.message.size, NextRetry = self .message .next_delivery_event(None) .map(trc::Value::Timestamp), NextDsn = self.message.next_dsn(None).map(trc::Value::Timestamp), Expires = self.message.expires(None).map(trc::Value::Timestamp), ); // Write message to queue let mut batch = BatchBuilder::new(); // Reserve quotas for quota_key in &self.message.quota_keys { match quota_key { QuotaKey::Count { key, .. } => { batch.add(ValueClass::Queue(QueueClass::QuotaCount(key.to_vec())), 1); } QuotaKey::Size { key, .. } => { batch.add( ValueClass::Queue(QueueClass::QuotaSize(key.to_vec())), self.message.size as i64, ); } } } for (queue_name, due) in self.message.next_events() { batch.set( ValueClass::Queue(QueueClass::MessageEvent(store::write::QueueEvent { due, queue_id: self.queue_id, queue_name: queue_name.into_inner(), })), Vec::new(), ); } if let Some((is_spam, subject)) = train_spam && let Some(config) = &server.core.spam.classifier { let hold_period = now + config.hold_samples_for; let sample = SpamTrainingSample { account_id: None, blob_id: BlobId::new(self.message.blob_hash.clone(), Default::default()), delete_after_use: false, expires_at: UTCDateTime::from_timestamp(hold_period as i64), from: self.message.return_path.to_string(), is_spam, subject, } .to_pickled_vec(); let object_id = ObjectType::SpamTrainingSample.to_id(); let item_id = SnowflakeIdGenerator::from_sequence_id(xxhash_rust::xxh3::xxh3_64( sample.as_slice(), )) .unwrap_or_default(); batch .set( BlobOp::Link { hash: self.message.blob_hash.clone(), to: BlobLink::Temporary { until: hold_period }, }, ObjectId::new(ObjectType::SpamTrainingSample, item_id.into()).serialize(), ) .set( ValueClass::Registry(RegistryClass::Item { object_id, item_id }), sample, ) .set( ValueClass::Registry(RegistryClass::Index { index_id: Property::AccountId.to_id(), object_id, item_id, key: (u32::MAX as u64).serialize(), }), vec![], ); trc::event!( Spam(SpamEvent::TrainSampleAdded), Details = if is_spam { "spam" } else { "ham" }, Expires = trc::Value::Timestamp(hold_period), SpanId = self.span_id, ); } batch .clear(BlobOp::Link { hash: self.message.blob_hash.clone(), to: BlobLink::Temporary { until: reserve_until, }, }) .set( BlobOp::Link { hash: self.message.blob_hash.clone(), to: BlobLink::Id { id: self.queue_id }, }, vec![], ) .set( BlobOp::Commit { hash: self.message.blob_hash.clone(), }, vec![], ) .set( ValueClass::Queue(QueueClass::Message(self.queue_id)), match Archiver::new(self.message).serialize() { Ok(data) => data, Err(err) => { trc::error!( err.details("Failed to serialize message.") .span_id(session_id) .caused_by(trc::location!()) ); return false; } }, ); if let Err(err) = server.store().write(batch.build_all()).await { trc::error!( err.details("Failed to write to store.") .span_id(session_id) .caused_by(trc::location!()) ); return false; } // Queue the message if server .inner .ipc .queue_tx .send(QueueEvent::Refresh) .await .is_err() { trc::event!( Server(ServerEvent::ThreadError), Reason = "Channel closed.", CausedBy = trc::location!(), SpanId = session_id, ); } true } pub async fn add_recipient(&mut self, rcpt: impl AsRef, server: &Server) { // Resolve queue self.message.recipients.push(Recipient::new(rcpt)); let queue = server.get_queue_or_default( &server .eval_if::( &server.core.smtp.queue.queue, &QueueEnvelope::new(&self.message, self.message.recipients.last().unwrap()), self.span_id, ) .await .unwrap_or_else(|| "default".to_string()), self.span_id, ); // Update expiration let now = now(); let recipient = self.message.recipients.last_mut().unwrap(); recipient.notify = Schedule::later(queue.notify.first().copied().unwrap_or(86400) + now); recipient.expires = queue.expiry; recipient.queue = queue.virtual_queue; } pub async fn save_changes(mut self, server: &Server, prev_event: Option) -> bool { // Release quota for completed deliveries let mut batch = BatchBuilder::new(); self.release_quota(&mut batch); // Update message queue if let Some(prev_event) = prev_event { batch.clear(ValueClass::Queue(QueueClass::MessageEvent( store::write::QueueEvent { due: prev_event, queue_id: self.queue_id, queue_name: self.queue_name.into_inner(), }, ))); } for (queue_name, due) in self.message.next_events() { batch.set( ValueClass::Queue(QueueClass::MessageEvent(store::write::QueueEvent { due, queue_id: self.queue_id, queue_name: queue_name.into_inner(), })), Vec::new(), ); } let message_bytes = match Archiver::new(self.message).serialize() { Ok(data) => data, Err(err) => { trc::error!( err.details("Failed to serialize message.") .span_id(self.span_id) .caused_by(trc::location!()) ); return false; } }; if self.is_multi_queue { batch.merge_fnc( ValueClass::Queue(QueueClass::Message(self.queue_id)), Params::with_capacity(3) .with_u64(self.queue_id) .with_bytes(self.queue_name.into_inner().to_vec()) .with_bytes(message_bytes), |params, _, bytes| { let mut cur_message = as Deserialize>::deserialize( bytes.ok_or_else(|| { trc::StoreEvent::NotFound .into_err() .details("Message no longer exists.") .caused_by(trc::location!()) .ctx(trc::Key::QueueId, params.u64(0)) })?, ) .and_then(|archive| archive.deserialize::()) .caused_by(trc::location!())?; let new_message_ = as Deserialize>::deserialize(params.bytes(2)) .caused_by(trc::location!())?; let new_message = new_message_ .unarchive::() .caused_by(trc::location!())?; if cur_message.blob_hash.as_slice() == new_message.blob_hash.0.as_slice() && cur_message.recipients.len() == new_message.recipients.len() { let queue_name = params.bytes(1); for (rcpt_idx, rcpt) in new_message .recipients .iter() .enumerate() .filter(|(_, rcpt)| rcpt.queue.as_slice() == queue_name) { cur_message.recipients[rcpt_idx] = rkyv_deserialize(rcpt).caused_by(trc::location!())?; } Archiver::new(cur_message) .serialize() .caused_by(trc::location!()) .map(MergeResult::Update) } else { Err(trc::StoreEvent::UnexpectedError .into_err() .details("Message blob hash or recipient count mismatch.") .caused_by(trc::location!()) .ctx(trc::Key::QueueId, params.u64(0))) } }, ); } else { batch.set( ValueClass::Queue(QueueClass::Message(self.queue_id)), message_bytes, ); } if let Err(err) = server.store().write(batch.build_all()).await { trc::error!( err.details("Failed to save changes.") .span_id(self.span_id) .caused_by(trc::location!()) ); false } else { true } } pub async fn remove(self, server: &Server, prev_event: Option) -> bool { let mut batch = BatchBuilder::new(); if let Some(prev_event) = prev_event { batch.clear(ValueClass::Queue(QueueClass::MessageEvent( store::write::QueueEvent { due: prev_event, queue_id: self.queue_id, queue_name: self.queue_name.into_inner(), }, ))); } else { for (queue_name, due) in self.message.next_events() { batch.clear(ValueClass::Queue(QueueClass::MessageEvent( store::write::QueueEvent { due, queue_id: self.queue_id, queue_name: queue_name.into_inner(), }, ))); } } // Release all quotas for quota_key in self.message.quota_keys { match quota_key { QuotaKey::Count { key, .. } => { batch.add(ValueClass::Queue(QueueClass::QuotaCount(key.to_vec())), -1); } QuotaKey::Size { key, .. } => { batch.add( ValueClass::Queue(QueueClass::QuotaSize(key.to_vec())), -(self.message.size as i64), ); } } } batch .clear(BlobOp::Link { hash: self.message.blob_hash.clone(), to: BlobLink::Id { id: self.queue_id }, }) .clear(ValueClass::Queue(QueueClass::Message(self.queue_id))); if let Err(err) = server.store().write(batch.build_all()).await { trc::error!( err.details("Failed to write to update queue.") .span_id(self.span_id) .caused_by(trc::location!()) ); false } else { true } } pub fn has_domain(&self, domains: &[String]) -> bool { self.message.recipients.iter().any(|r| { let domain = r.address.domain_part(); domains.iter().any(|dd| dd == domain) }) || self .message .return_path .rsplit_once('@') .is_some_and(|(_, domain)| domains.iter().any(|dd| dd == domain)) } } impl ArchivedMessage { pub fn has_domain(&self, domains: &[String]) -> bool { self.recipients.iter().any(|r| { let domain = r.address.domain_part(); domains.iter().any(|dd| dd == domain) }) || self .return_path .rsplit_once('@') .is_some_and(|(_, domain)| domains.iter().any(|dd| dd == domain)) } pub fn next_delivery_event(&self, queue: Option) -> Option { let mut next_delivery = None; for rcpt in self.recipients.iter().filter(|d| { matches!( d.status, ArchivedStatus::Scheduled | ArchivedStatus::TemporaryFailure(_) ) && queue.is_none_or(|q| d.queue == q) }) { let retry_due = rcpt.retry.due.to_native(); if let Some(next_delivery) = &mut next_delivery { if retry_due < *next_delivery { *next_delivery = retry_due; } } else { next_delivery = Some(retry_due); } } next_delivery } }