use std::io::{self, Read, Write};
use std::mem;
use openssl::ssl::{SslAcceptor, SslConnector, SslMethod, SslVerifyMode};
use crate::test_data::{CERTIFICATE, CERTIFICATE_PRIVATE_KEY};
pub fn ssl_acceptor() -> SslAcceptor {
let mut ssl_acceptor =
SslAcceptor::mozilla_intermediate_v5(SslMethod::tls_server()).unwrap();
ssl_acceptor
.set_private_key(&CERTIFICATE_PRIVATE_KEY)
.unwrap();
ssl_acceptor.set_certificate(&CERTIFICATE).unwrap();
ssl_acceptor.build()
}
pub trait ReadWrite: Read + Write {}
impl<T: Read + Write + ?Sized> ReadWrite for T {}
pub struct SmtpClient {
name: &'static str,
io: Box<dyn ReadWrite>,
}
impl SmtpClient {
pub fn new(name: &'static str, io: impl ReadWrite + 'static) -> Self {
Self {
name,
io: Box::new(io),
}
}
pub fn read_responses(&mut self) -> Vec<String> {
let mut ret = Vec::<String>::new();
loop {
let mut line_bytes = Vec::<u8>::new();
while Some(b'\n') != line_bytes.last().copied() {
let mut buf = [0u8; 1];
let nread = self.io.read(&mut buf).unwrap();
if 0 == nread {
panic!("Unexpected EOF");
}
line_bytes.push(buf[0]);
}
let line = String::from_utf8(line_bytes).unwrap();
let last = " " == &line[3..4];
println!("[{}] >> {:?}", self.name, line);
ret.push(line);
if last {
break;
}
}
ret
}
pub fn write_line(&mut self, s: &str) {
assert!(s.ends_with('\n'));
for line in s.split_inclusive('\n') {
println!("[{}] << {:?}", self.name, line);
}
self.io.write_all(s.as_bytes()).unwrap();
}
pub fn write_raw(&mut self, data: &[u8]) {
println!("[{}] << [{} bytes]", self.name, data.len());
self.io.write_all(data).unwrap();
}
pub fn skip_pleasantries(&mut self, cmd: &str) {
self.read_responses();
self.write_line(&format!("{}\r\n", cmd));
let responses = self.read_responses();
assert!(responses.last().unwrap().starts_with("250"));
}
pub fn simple_command(&mut self, command: &str, prefix: &str) {
self.write_line(&format!("{}\r\n", command));
let responses = self.read_responses();
assert_eq!(1, responses.len());
assert!(responses[0].starts_with(prefix));
}
pub fn unix_simple_command(&mut self, command: &str, prefix: &str) {
self.write_line(&format!("{}\n", command));
let responses = self.read_responses();
assert_eq!(1, responses.len());
assert!(responses[0].starts_with(prefix));
}
pub fn start_tls(&mut self) {
let mut connector = SslConnector::builder(SslMethod::tls()).unwrap();
connector.set_verify(SslVerifyMode::NONE);
println!("[{}] <> Start TLS handshake", self.name);
let cxn = mem::replace(&mut self.io, Box::new(io::empty()));
let cxn = connector
.build()
.connect("localhost", cxn)
.map_err(|_| "SSL handshake failed")
.unwrap();
println!("[{}] <> TLS handshake succeeded", self.name);
self.io = Box::new(cxn);
}
pub fn skip_pleasantries_with_tls(&mut self, command: &str) {
self.skip_pleasantries(command);
self.simple_command("STARTTLS", "220 2.0.0");
self.start_tls();
self.write_line(&format!("{}\r\n", command));
let responses = self.read_responses();
assert!(responses.last().unwrap().starts_with("250"));
}
pub fn quick_log_in(&mut self, helo: &str, user: &str, password: &str) {
self.skip_pleasantries_with_tls(helo);
let auth = format!(
"AUTH PLAIN {}",
base64::encode(format!("{user}\0{user}\0{password}")),
);
self.simple_command(&auth, "235 ");
}
pub fn assert_eof(&mut self) {
let mut buf = [0u8; 1];
assert_matches!(
io::ErrorKind::UnexpectedEof | io::ErrorKind::BrokenPipe,
self.io.read_exact(&mut buf).unwrap_err().kind(),
);
}
}