use crate::{
common::{Arc, Dummy},
io::ConnectionInfo,
mail::MailService,
smtp::SmtpSession,
};
use std::{
any::{Any, TypeId},
collections::HashMap,
};
#[derive(Debug, Default)]
pub struct SmtpContext {
store: HashMap<TypeId, Box<dyn Any + Send + Sync>>,
pub session: SmtpSession,
}
impl SmtpContext {
pub fn new<Svc>(service: Svc, connection: ConnectionInfo) -> Self
where
Svc: MailService + Sync + Send + 'static,
{
let mut me = SmtpContext {
session: SmtpSession::new(connection),
..Default::default()
};
me.set_service(service);
me
}
}
impl SmtpContext {
pub fn emtpy() -> Self {
Self::default()
}
pub fn get<T: Sync + Send + 'static>(&self) -> Option<&T> {
self.store
.get(&TypeId::of::<T>())
.and_then(|v| v.downcast_ref())
}
pub fn get_mut<T: Sync + Send + 'static>(&mut self) -> Option<&mut T> {
self.store
.get_mut(&TypeId::of::<T>())
.and_then(|v| v.downcast_mut())
}
pub fn get_or_insert<T: Sync + Send + 'static, F>(&mut self, insert: F) -> &mut T
where
F: FnOnce() -> T,
{
let id = TypeId::of::<T>();
self.store
.entry(id)
.or_insert_with(|| Box::new(insert()))
.downcast_mut::<T>()
.expect("stored type must match")
}
pub fn set<T: Sync + Send + 'static>(&mut self, value: T) {
let id = TypeId::of::<T>();
self.store.insert(id, Box::new(value));
}
pub fn service(&self) -> impl MailService {
self.get::<Arc<dyn MailService + Send + Sync + 'static>>()
.cloned()
.unwrap_or_else(|| Arc::new(Dummy) as Arc<dyn MailService + Send + Sync>)
}
pub(crate) fn set_service(&mut self, service: impl MailService + Send + Sync + 'static) {
let service = Arc::new(service) as Arc<dyn MailService + Send + Sync + 'static>;
self.set(service);
}
}
#[derive(Clone, Eq, PartialEq)]
pub enum DriverControl {
Response(Vec<u8>),
StartTls,
Shutdown,
}
impl std::fmt::Debug for DriverControl {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
#[derive(Debug)]
enum TextOrBytes<'a> {
T(&'a str),
B(&'a [u8]),
}
fn tb(inp: &[u8]) -> TextOrBytes {
if let Ok(text) = std::str::from_utf8(inp) {
TextOrBytes::T(text)
} else {
TextOrBytes::B(inp)
}
}
match self {
DriverControl::Response(r) => f.debug_tuple("Response").field(&tb(r)).finish(),
DriverControl::StartTls => f.debug_tuple("StartTls").finish(),
DriverControl::Shutdown => f.debug_tuple("Shutdown").finish(),
}
}
}
#[cfg(test)]
mod store_tests {
use super::*;
use crate::mail::SessionLogger;
use regex::Regex;
#[test]
pub fn same_service() {
let mut sut = SmtpContext::default();
let svc = Box::new(SessionLogger);
let dump0 = format!("{:#?}", svc);
sut.set_service(Dummy);
sut.set_service(svc);
let dump1 = format!("{:#?}", sut.service());
assert_eq!(dump1, dump0);
insta::assert_display_snapshot!(dump0, @r###"SessionLogger"###);
}
#[test]
pub fn set_one_service() {
let mut sut = SmtpContext::default();
sut.set_service(Box::new(Dummy));
sut.set_service(SessionLogger);
let dump = format!("{:#?}", sut);
let dump = Regex::new("[0-9]+")
.expect("regex")
.replace_all(dump.as_str(), "--redacted--");
insta::assert_display_snapshot!(dump, @r###"
SmtpContext {
store: {
TypeId {
t: --redacted--,
}: Any { .. },
},
session: SmtpSession {
connection: ConnectionInfo {
id: "--redacted--",
local_addr: "",
peer_addr: "",
established: SystemTime {
tv_sec: --redacted--,
tv_nsec: --redacted--,
},
},
extensions: ExtensionSet {
map: {},
},
service_name: "samotop",
peer_name: None,
output: [],
input: [],
mode: None,
transaction: Transaction {
id: "",
mail: None,
rcpts: [],
extra_headers: "",
sink: "*",
},
},
}
"###);
}
}