use crate::utils::network::tls_config::get_client_config;
use rustls::{ClientConnection, ServerName, StreamOwned};
use std::io;
use std::io::Read;
use std::io::Write;
use std::net::TcpStream;
use std::sync::Arc;
#[derive(Debug)]
pub enum TLSConnectError {
Rustls(rustls::Error),
InvalidDNS(rustls::client::InvalidDnsNameError),
Io(io::Error),
}
pub trait TLSStreamClient {
fn connect(address: &(String, u16)) -> Result<Self, TLSConnectError>
where
Self: Sized;
fn get_tlsstream_mut(&mut self) -> &mut StreamOwned<ClientConnection, TcpStream>;
fn read(&mut self) -> Result<String, io::Error> {
let mut result_string = String::new();
let buffer_size = 512;
let mut continue_reading = true;
while continue_reading {
let mut buf = vec![0; buffer_size];
let bytes = self.get_tlsstream_mut().read(&mut buf)?;
buf.truncate(bytes);
result_string.push_str(
&String::from_utf8(buf).map_err(|e| io::Error::new(io::ErrorKind::Other, e))?,
);
continue_reading = bytes == buffer_size && bytes > 0;
}
Ok(result_string)
}
fn write(&mut self, data: String) -> Result<(), io::Error> {
let buf = data.as_bytes();
self.get_tlsstream_mut().write_all(buf)?;
self.get_tlsstream_mut().flush()?;
Ok(())
}
}
pub struct SimpleTLSStreamClient {
tls_stream: StreamOwned<ClientConnection, TcpStream>,
}
impl TLSStreamClient for SimpleTLSStreamClient {
fn connect(address: &(String, u16)) -> Result<Self, TLSConnectError>
where
Self: Sized,
{
let tcp_stream = TcpStream::connect(address).map_err(TLSConnectError::Io)?;
let server_name =
ServerName::try_from(address.0.as_str()).map_err(TLSConnectError::InvalidDNS)?;
let config = Arc::new(get_client_config().map_err(TLSConnectError::Rustls)?);
let session =
ClientConnection::new(config, server_name).map_err(TLSConnectError::Rustls)?;
let tls_stream = StreamOwned::new(session, tcp_stream);
Ok(Self { tls_stream })
}
fn get_tlsstream_mut(&mut self) -> &mut StreamOwned<ClientConnection, TcpStream> {
&mut self.tls_stream
}
}