use std::time::Duration;
use std::{net::SocketAddr, sync::Arc};
use tokio::io::BufStream;
use tokio::sync::Mutex;
use tokio::time::timeout;
use tokio_native_tls::TlsAcceptor;
use crate::client_message::ClientMessage;
use crate::controllers::on_auth::OnAuthController;
use crate::mail::Mail;
use super::command::Commands;
use super::connection::SMTPConnection;
use super::connection::SMTPConnectionStatus;
use super::controllers::on_close::OnCloseController;
use super::controllers::on_email::OnEmailController;
use super::controllers::on_reset::OnResetController;
use super::errors::SMTPError;
use super::message::Message;
use super::status_code::StatusCodes;
#[derive(Debug)]
pub struct SMTPServer<B> {
use_tls: bool,
listener: Option<Arc<tokio::net::TcpListener>>,
workers: usize,
threads_pool: Option<Arc<rayon::ThreadPool>>,
tls_acceptor: Option<Arc<Mutex<tokio_native_tls::TlsAcceptor>>>,
controllers: Controllers<B>,
max_size: usize,
allowed_commands: Vec<Commands>,
max_session_duration: Duration,
max_op_duration: Duration,
}
#[derive(Debug)]
pub struct Controllers<B> {
on_auth: Option<OnAuthController<B>>,
on_email: Option<OnEmailController<B>>,
on_reset: Option<OnResetController<B>>,
on_close: Option<OnCloseController<B>>,
}
impl<B> Clone for Controllers<B>
where
B: Default + Send + Sync + Clone,
{
fn clone(&self) -> Self {
Controllers {
on_auth: self.on_auth.clone(),
on_email: self.on_email.clone(),
on_reset: self.on_reset.clone(),
on_close: self.on_close.clone(),
}
}
}
impl<B> SMTPServer<B> {
pub fn new() -> Self {
SMTPServer {
use_tls: false,
listener: None,
workers: 1,
threads_pool: None,
tls_acceptor: None,
controllers: Controllers {
on_auth: None,
on_email: None,
on_reset: None,
on_close: None,
},
max_size: 1024 * 1024 * 10, allowed_commands: vec![
Commands::HELO,
Commands::EHLO,
Commands::MAIL,
Commands::RCPT,
Commands::DATA,
Commands::RSET,
Commands::VRFY,
Commands::EXPN,
Commands::HELP,
Commands::NOOP,
Commands::QUIT,
Commands::AUTH,
Commands::STARTTLS,
],
max_session_duration: Duration::from_secs(300),
max_op_duration: Duration::from_secs(30),
}
}
pub fn workers(&mut self, workers: usize) -> &mut Self {
log::info!("Setting workers to {}", workers);
self.workers = workers;
self
}
pub fn set_tls_acceptor(&mut self, acceptor: tokio_native_tls::TlsAcceptor) -> &mut Self {
log::info!("TLS Acceptor set");
self.use_tls = true;
self.tls_acceptor = Some(Arc::new(Mutex::new(acceptor)));
self
}
pub fn set_max_size(&mut self, max_size: usize) -> &mut Self {
log::info!("Setting max size to {}", max_size);
self.max_size = max_size;
self
}
pub fn set_allowed_commands(&mut self, commands: Vec<Commands>) -> &mut Self {
log::info!("Setting allowed commands");
self.allowed_commands = commands;
self
}
pub fn on_auth(&mut self, on_auth: OnAuthController<B>) -> &mut Self {
log::info!("Setting OnAuthController");
self.controllers.on_auth = Some(on_auth);
self
}
pub fn on_email(&mut self, on_email: OnEmailController<B>) -> &mut Self {
log::info!("Setting OnEmailController");
self.controllers.on_email = Some(on_email);
self
}
pub fn on_reset(&mut self, on_reset: OnResetController<B>) -> &mut Self {
log::info!("Setting OnResetController");
self.controllers.on_reset = Some(on_reset);
self
}
pub fn on_close(&mut self, on_close: OnCloseController<B>) -> &mut Self {
log::info!("Setting OnCloseController");
self.controllers.on_close = Some(on_close);
self
}
pub fn set_max_session_duration(&mut self, duration: Duration) -> &mut Self {
log::info!("Setting max session duration to {:?}", duration);
self.max_session_duration = duration;
self
}
pub fn set_max_op_duration(&mut self, duration: Duration) -> &mut Self {
log::info!("Setting max operation duration to {:?}", duration);
self.max_op_duration = duration;
self
}
pub async fn bind(&mut self, address: SocketAddr) -> Result<&mut Self, tokio::io::Error> {
log::info!("Binding to {}", address);
let listener = tokio::net::TcpListener::bind(address).await?;
self.listener = Some(Arc::new(listener));
Ok(self)
}
pub async fn run(&mut self)
where
B: 'static + Default + Send + Sync + Clone,
{
let listener = match self.listener.clone() {
Some(lstnr) => lstnr,
None => panic!("There isn't listener"),
};
log::info!("Building ThreadPool with {} workers", self.workers);
self.threads_pool = match rayon::ThreadPoolBuilder::new()
.num_threads(self.workers)
.build()
{
Ok(pool) => Some(Arc::new(pool)),
Err(err) => panic!("{}", err),
};
log::info!("Starting main loop for accepting connections");
loop {
let (socket, _) = match listener.accept().await {
Ok(conn) => conn,
Err(err) => {
log::error!(
"An error ocurred while trying to accept and TcpStream connection {}",
err
);
continue;
}
};
log::debug!("Connection received from {}", socket.peer_addr().unwrap());
let pool = self.threads_pool.clone();
let use_tls = self.use_tls;
let tls_acceptor = self.tls_acceptor.clone();
let controllers = self.controllers.clone();
let max_size = self.max_size;
let allowed_commands = self.allowed_commands.clone();
let max_session_duration = self.max_session_duration;
let max_op_duration = self.max_op_duration;
tokio::spawn(async move {
log::debug!("Initializing TCP connection");
let conn = Arc::new(Mutex::new(SMTPConnection {
use_tls: false,
tls_buff_socket: None,
tcp_buff_socket: Some(Arc::new(Mutex::new(BufStream::new(socket)))),
buffer: Vec::new(),
mail_buffer: Vec::new(),
status: SMTPConnectionStatus::WaitingCommand,
state: Arc::new(Mutex::new(B::default())),
}));
if let Some(pool) = pool {
pool.install(|| {
tokio::runtime::Runtime::new().unwrap().block_on(
handle_connection_with_timeout(
use_tls,
tls_acceptor,
conn,
controllers,
max_size,
allowed_commands,
max_session_duration,
max_op_duration,
),
);
});
}
});
}
}
}
pub async fn handle_connection_with_timeout<B>(
use_tls: bool,
tls_acceptor: Option<Arc<Mutex<TlsAcceptor>>>,
mutex_con: Arc<Mutex<SMTPConnection<B>>>,
controllers: Controllers<B>,
max_size: usize,
allowed_commands: Vec<Commands>,
max_session_duration: Duration,
max_op_duration: Duration,
) where
B: 'static + Default + Send + Sync + Clone,
{
let mutex_conn_for_handle_connection = mutex_con.clone();
match timeout(
max_session_duration,
handle_connection(
use_tls,
tls_acceptor,
mutex_conn_for_handle_connection,
controllers,
max_size,
allowed_commands,
max_op_duration,
),
)
.await
{
Ok(_) => (),
Err(_) => {
let conn = mutex_con.lock().await;
let _ = conn
.write_socket(
&Message::builder()
.status(StatusCodes::ServiceClosingTransmissionChannel)
.message("Service closing transmission channel".to_string())
.build()
.as_bytes(true),
)
.await
.map_err(|err| log::error!("{}", err));
let _ = conn.close().await.map_err(|err| log::error!("{}", err));
}
}
}
pub async fn handle_connection<B>(
use_tls: bool,
tls_acceptor: Option<Arc<Mutex<TlsAcceptor>>>,
mutex_con: Arc<Mutex<SMTPConnection<B>>>,
controllers: Controllers<B>,
max_size: usize,
allowed_commands: Vec<Commands>,
max_op_duration: Duration,
) where
B: 'static + Default + Send + Sync + Clone,
{
log::trace!("[] Handling connection with optional TLS?: {}", use_tls);
let conn = mutex_con.lock().await;
match conn
.write_socket(
&Message::builder()
.status(StatusCodes::SMTPServiceReady)
.message("SMTP Service Ready".to_string())
.build()
.as_bytes(true),
)
.await
{
Ok(_) => (),
Err(err) => panic!("{}", err),
};
drop(conn);
log::trace!("[] Connection initialized, and start proccessing commands");
loop {
match timeout(
max_op_duration,
handle_connection_logic(
use_tls,
tls_acceptor.clone(),
mutex_con.clone(),
controllers.clone(),
max_size,
allowed_commands.clone(),
),
)
.await
{
Ok(HandleConnectionFlow::Continue) => (),
Ok(HandleConnectionFlow::Break) => break,
Err(_) => {
log::trace!("[❌] Timeout reached, closing connection");
break;
}
}
}
let conn = mutex_con.lock().await;
controllers.on_close.as_ref().map(|on_close| {
let on_close = on_close.0.clone();
drop(conn);
let _ = on_close(mutex_con.clone());
});
let conn = mutex_con.lock().await;
log::trace!("[❌] Sending final message to client to close");
let _ = conn
.write_socket(
&Message::builder()
.status(StatusCodes::ServiceClosingTransmissionChannel)
.message("Service closing transmission channel".to_string())
.build()
.as_bytes(true),
)
.await
.map_err(|err| log::error!("{}", err));
log::trace!("[❌] Closing connection with client");
let _ = conn.close().await.map_err(|err| log::error!("{}", err));
}
pub enum HandleConnectionFlow {
Continue,
Break,
}
pub async fn handle_connection_logic<B>(
use_tls: bool,
tls_acceptor: Option<Arc<Mutex<TlsAcceptor>>>,
mutex_con: Arc<Mutex<SMTPConnection<B>>>,
controllers: Controllers<B>,
max_size: usize,
allowed_commands: Vec<Commands>,
) -> HandleConnectionFlow
where
B: 'static + Default + Send + Sync + Clone,
{
let mut conn = mutex_con.lock().await;
let mut buf = [0; 2048];
let n = conn.read_socket(&mut buf).await.unwrap_or_else(|err| {
log::trace!("[🕵️♂️💻] Error reading from socket: {}", err);
0
});
if n == 0 {
drop(conn);
log::trace!("[🖥️🔒] Connection closed by client");
return HandleConnectionFlow::Break;
}
if conn.status == SMTPConnectionStatus::WaitingCommand && conn.buffer.len() + n > 2048 {
let _ = conn
.write_socket(
&Message::builder()
.status(StatusCodes::ExceededStorageAllocation)
.message("Buffer size exceeded, Resetting buffer".to_string())
.build()
.as_bytes(true),
)
.await
.map_err(|err| log::error!("{}", err));
conn.buffer.clear();
controllers.on_reset.as_ref().map(|on_reset| {
let on_reset = on_reset.0.clone();
drop(conn);
let _ = on_reset(mutex_con.clone());
});
return HandleConnectionFlow::Continue;
}
if conn.status == SMTPConnectionStatus::WaitingData && conn.mail_buffer.len() + n > max_size {
let _ = conn
.write_socket(
&Message::builder()
.status(StatusCodes::ExceededStorageAllocation)
.message("Buffer size exceeded, Resetting buffer".to_string())
.build()
.as_bytes(true),
)
.await
.map_err(|err| log::error!("{}", err));
conn.mail_buffer.clear();
controllers.on_reset.as_ref().map(|on_reset| {
let on_reset = on_reset.0.clone();
drop(conn);
let _ = on_reset(mutex_con.clone());
});
return HandleConnectionFlow::Continue;
}
if conn.status == SMTPConnectionStatus::WaitingData {
conn.mail_buffer.extend_from_slice(&buf[..n]);
} else {
conn.buffer.extend_from_slice(&buf[..n]);
}
if conn.status == SMTPConnectionStatus::WaitingData && conn.mail_buffer.ends_with(b"\r\n.\r\n")
{
if let Some(on_email) = &controllers.on_email {
let on_email = on_email.0.clone();
let mail = match Mail::<Vec<u8>>::from_bytes(conn.mail_buffer.clone()) {
Ok(mail) => mail,
Err(err) => {
log::error!("{}", err);
return HandleConnectionFlow::Continue;
}
};
drop(conn);
let response = on_email(mutex_con.clone(), Box::new(mail)).await;
let conn = mutex_con.lock().await;
let _ = conn
.write_socket(&response.as_bytes(true))
.await
.map_err(|err| {
log::error!("{}", err);
});
} else {
let response = Message::builder()
.status(StatusCodes::OK)
.message("Message received".to_string())
.build()
.to_string(true);
conn.write_socket(response.as_bytes()).await.unwrap();
}
let mut conn = mutex_con.lock().await;
conn.mail_buffer.clear();
conn.buffer.clear();
conn.status = SMTPConnectionStatus::WaitingCommand;
return HandleConnectionFlow::Continue;
}
if conn.status == SMTPConnectionStatus::WaitingCommand && conn.buffer.ends_with(b"\r\n") {
let mut client_message = match ClientMessage::<String>::from_bytes(conn.buffer.clone()) {
Ok(msg) => msg,
Err(err) => {
match conn
.write_socket(
&Message::builder()
.status(StatusCodes::SyntaxError)
.message(err.to_string())
.build()
.as_bytes(true),
)
.await
{
Ok(_) => (),
Err(err) => {
log::error!("{}", err);
return HandleConnectionFlow::Continue;
}
}
return HandleConnectionFlow::Continue;
}
};
if client_message.command == Commands::QUIT {
return HandleConnectionFlow::Break;
} else if client_message.command == Commands::RSET {
conn.buffer.clear();
conn.status = SMTPConnectionStatus::WaitingCommand;
controllers.on_reset.as_ref().map(|on_reset| {
let on_reset = on_reset.0.clone();
drop(conn);
let _ = on_reset(mutex_con.clone());
});
return HandleConnectionFlow::Continue;
}
log::trace!("[💬] Received Message: {:?}", client_message);
drop(conn);
let (mut response, status) = match handle_command(
mutex_con.clone(),
controllers.clone(),
&mut client_message,
allowed_commands.clone(),
max_size,
)
.await
{
Ok((res, status)) => (res, status),
Err(err) => {
let conn = mutex_con.lock().await;
let _ = conn
.write_socket(
&Message::builder()
.status(StatusCodes::TransactionFailed)
.message(err.to_string())
.build()
.as_bytes(true),
)
.await
.map_err(|err| log::error!("{}", err));
return HandleConnectionFlow::Continue;
}
};
log::trace!(
"[💬] Response for SMTP command {:?} is: {:?}",
client_message.command,
response
);
let mut conn = mutex_con.lock().await;
conn.status = status;
let last_index = response.len() - 1;
let tls_acceptor = tls_acceptor.clone();
if conn.status == SMTPConnectionStatus::Closed {
for (i, message) in response.iter_mut().enumerate() {
let is_last = i == last_index;
let bytes = message.as_bytes(is_last);
conn.write_socket(&bytes).await.unwrap();
}
conn.buffer.clear();
return HandleConnectionFlow::Break;
} else if conn.status == SMTPConnectionStatus::StartTLS && use_tls && tls_acceptor.is_some() {
match conn
.write_socket(
&Message::builder()
.status(StatusCodes::SMTPServiceReady)
.message("Ready to start TLS".to_string())
.build()
.as_bytes(true),
)
.await
{
Ok(_) => (),
Err(err) => {
log::error!("{}", err);
return HandleConnectionFlow::Break;
}
}
log::debug!("Upgrading connection to TLS");
drop(conn);
match upgrade_to_tls(mutex_con.clone(), tls_acceptor).await {
Ok(_) => {
log::debug!("Connection upgraded to TLS");
let mut conn = mutex_con.lock().await;
conn.buffer.clear();
conn.status = SMTPConnectionStatus::WaitingCommand;
return HandleConnectionFlow::Continue;
}
Err(err) => {
log::error!("An error ocurred while trying to upgrade to TLS {}", err);
let mut conn = mutex_con.lock().await;
conn.write_socket(
&Message::builder()
.status(StatusCodes::TransactionFailed)
.message("TLS not available".to_string())
.build()
.as_bytes(true),
)
.await
.unwrap();
conn.buffer.clear();
conn.status = SMTPConnectionStatus::WaitingCommand;
}
};
} else if conn.status == SMTPConnectionStatus::StartTLS && !use_tls {
log::trace!("[🛡️] TLS not available");
let _ = conn
.write_socket(
&Message::builder()
.status(StatusCodes::TransactionFailed)
.message("TLS not available".to_string())
.build()
.as_bytes(true),
)
.await
.map_err(|err| log::error!("{}", err));
conn.buffer.clear();
conn.status = SMTPConnectionStatus::WaitingCommand;
} else {
for (i, message) in response.iter_mut().enumerate() {
let is_last = i == last_index;
let bytes = message.as_bytes(is_last);
conn.write_socket(&bytes).await.unwrap();
}
conn.buffer.clear();
}
}
HandleConnectionFlow::Continue
}
pub async fn upgrade_to_tls<B>(
conn: Arc<Mutex<SMTPConnection<B>>>,
tls_acceptor: Option<Arc<Mutex<tokio_native_tls::TlsAcceptor>>>,
) -> Result<(), Box<dyn std::error::Error>> {
log::trace!("[🌐🔒] Upgrading connection to TLS");
let tls_acceptor = match tls_acceptor {
Some(tls_acceptor) => tls_acceptor,
None => return Err("TLS Acceptor not set".into()),
};
log::trace!("[🌐🔒] Locking connection to upgrade to TLS");
let mut conn_locked = conn.lock().await;
log::trace!("[🌐🔒] Connection locked");
let tcp_buff_socket = conn_locked
.tcp_buff_socket
.take()
.ok_or("No TcpStream found")?;
let tcp_buff_socket = Arc::try_unwrap(tcp_buff_socket).map_err(|_| "Failed to unwrap Arc")?;
let tcp_buff_socket = tcp_buff_socket.into_inner();
let tcp_stream = tcp_buff_socket.into_inner();
log::trace!("[🌐🔒] Locking TLS Acceptor");
let tls_acceptor = tls_acceptor.lock().await.clone();
log::trace!("[🌐🔒] TLS Acceptor locked");
log::trace!("[🌐🔒] Accepting TLS connection");
let tls_stream = match timeout(Duration::from_secs(10), tls_acceptor.accept(tcp_stream)).await {
Ok(Ok(tls_stream)) => {
log::trace!("[🌐🔒] TLS connection Accepted");
tls_stream
}
Ok(Err(err)) => {
log::error!("[🌐🔒] Error during TLS handshake: {}", err);
return Err(err.into());
}
Err(_) => {
log::error!("[🌐🔒] TLS handshake timed out");
return Err("TLS handshake timed out".into());
}
};
conn_locked.tls_buff_socket = Some(Arc::new(Mutex::new(BufStream::new(tls_stream))));
conn_locked.use_tls = true;
conn_locked.status = SMTPConnectionStatus::WaitingCommand;
Ok(())
}
pub async fn handle_command<B>(
conn: Arc<Mutex<SMTPConnection<B>>>,
controllers: Controllers<B>,
client_message: &mut ClientMessage<String>,
allowed_commands: Vec<Commands>,
max_size: usize,
) -> Result<(Vec<Message>, SMTPConnectionStatus), SMTPError>
where
B: 'static + Default + Send + Sync + Clone,
{
log::trace!("[⚙️] Handling SMTP command: {:?}", client_message.command);
if allowed_commands
.iter()
.find(|&cmd| cmd == &client_message.command)
.is_none()
{
return Err(SMTPError::UnknownCommand(client_message.command.clone()));
}
let result = match client_message.command {
Commands::HELO => (
vec![Message::builder()
.status(StatusCodes::OK)
.message(format!("Hello {}", "unknown"))
.build()],
SMTPConnectionStatus::WaitingCommand,
),
Commands::EHLO => {
let mut ehlo_messages = vec![
Message::builder()
.status(StatusCodes::OK)
.message("Hello".to_string())
.build(),
Message::builder()
.status(StatusCodes::OK)
.message(format!("SIZE {}", max_size))
.build(),
Message::builder()
.status(StatusCodes::OK)
.message("8BITMIME".to_string())
.build(),
Message::builder()
.status(StatusCodes::OK)
.message("PIPELINING".to_string())
.build(),
Message::builder()
.status(StatusCodes::OK)
.message("HELP".to_string())
.build(),
];
let conn = conn.lock().await;
if !conn.use_tls {
ehlo_messages.push(
Message::builder()
.status(StatusCodes::OK)
.message("STARTTLS".to_string())
.build(),
)
}
if controllers.on_auth.is_some() {
ehlo_messages.push(
Message::builder()
.status(StatusCodes::OK)
.message("AUTH PLAIN LOGIN CRAM-MD5 DIGEST-MD5 GSSAPI NTLM XOAUTH2".to_string())
.build(),
);
}
drop(conn);
(ehlo_messages, SMTPConnectionStatus::WaitingCommand)
}
Commands::MAIL => (
vec![Message::builder()
.status(StatusCodes::OK)
.message("Hello".to_string())
.build()],
SMTPConnectionStatus::WaitingCommand,
),
Commands::RCPT => (
vec![Message::builder()
.status(StatusCodes::OK)
.message("Hello".to_string())
.build()],
SMTPConnectionStatus::WaitingCommand,
),
Commands::DATA => (
vec![Message::builder()
.status(StatusCodes::StartMailInput)
.message("Start mail input; end with <CRLF>.<CRLF>".to_string())
.build()],
SMTPConnectionStatus::WaitingData,
),
Commands::RSET => (
vec![Message::builder()
.status(StatusCodes::OK)
.message("Hello".to_string())
.build()],
SMTPConnectionStatus::WaitingCommand,
),
Commands::VRFY => (
vec![Message::builder()
.status(StatusCodes::CannotVerifyUserButWillAcceptMessageAndAttemptDelivery)
.message(
"Cannot VRFY user, but will accept message and attempt delivery".to_string(),
)
.build()],
SMTPConnectionStatus::WaitingCommand,
),
Commands::EXPN => (
vec![Message::builder()
.status(StatusCodes::CommandNotImplemented)
.message(
"Cannot EXPN user, but will accept message and attempt delivery".to_string(),
)
.build()],
SMTPConnectionStatus::WaitingCommand,
),
Commands::HELP => (
vec![Message::builder()
.status(StatusCodes::HelpMessage)
.message("Help message".to_string())
.build()],
SMTPConnectionStatus::WaitingCommand,
),
Commands::NOOP => (
vec![Message::builder()
.status(StatusCodes::OK)
.message("NOOP Command successful".to_string())
.build()],
SMTPConnectionStatus::WaitingCommand,
),
Commands::QUIT => (
vec![Message::builder()
.status(StatusCodes::ServiceClosingTransmissionChannel)
.message("Service closing transmission channel".to_string())
.build()],
SMTPConnectionStatus::Closed,
),
Commands::AUTH => {
if let Some(on_auth) = &controllers.on_auth {
let on_auth = on_auth.0.clone();
match on_auth(conn.clone(), client_message.data.clone()).await {
Ok(response) => return Ok((vec![response], SMTPConnectionStatus::WaitingCommand)),
Err(response) => return Ok((vec![response], SMTPConnectionStatus::Closed)),
}
} else {
(
vec![Message::builder()
.status(StatusCodes::CommandNotImplemented)
.message("Command not recognized".to_string())
.build()],
SMTPConnectionStatus::WaitingCommand,
)
}
}
Commands::STARTTLS => (
vec![Message::builder()
.status(StatusCodes::SMTPServiceReady)
.message("Ready to start TLS".to_string())
.build()],
SMTPConnectionStatus::StartTLS,
),
_ => (
vec![Message::builder()
.status(StatusCodes::CommandNotImplemented)
.message("Command not recognized".to_string())
.build()],
SMTPConnectionStatus::WaitingCommand,
),
};
Ok(result)
}