use std::{
fs::File,
io::{self, BufReader},
net::SocketAddr,
sync::Arc,
};
use rustls::{
ClientConfig, RootCertStore,
pki_types::{CertificateDer, PrivateKeyDer},
};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpStream,
};
use tokio_rustls::TlsConnector;
#[derive(Debug)]
pub struct RustlsClient {
tls_config: Arc<ClientConfig>,
server_addr: SocketAddr,
}
impl RustlsClient {
pub fn new(cert_path: &str, key_path: &str, server_addr: SocketAddr) -> io::Result<Self> {
let cert_file = File::open(cert_path)?;
let mut cert_reader = BufReader::new(cert_file);
let certs = load_certs(&mut cert_reader)?;
let key_file = File::open(key_path)?;
let mut key_reader = BufReader::new(key_file);
let key = load_private_key(&mut key_reader)?;
let mut root_store = RootCertStore::empty();
root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
let config = ClientConfig::builder()
.dangerous() .with_custom_certificate_verifier(Arc::new(NoVerifier))
.with_client_auth_cert(certs, key)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?;
Ok(Self {
tls_config: Arc::new(config),
server_addr,
})
}
pub async fn connect(&self) -> io::Result<RustlsConnection> {
let tcp_stream = TcpStream::connect(&self.server_addr).await?;
let ip_addr = self.server_addr.ip();
let server_name = rustls::pki_types::ServerName::IpAddress(ip_addr.into());
let connector = TlsConnector::from(self.tls_config.clone());
let server_name = rustls::pki_types::ServerName::try_from(server_name)
.map_err(|e| io::Error::new(io::ErrorKind::InvalidInput, e))?
.to_owned();
let tls_stream = connector.connect(server_name, tcp_stream).await?;
Ok(RustlsConnection { stream: tls_stream })
}
}
#[derive(Debug)]
pub struct RustlsConnection {
stream: tokio_rustls::client::TlsStream<TcpStream>,
}
impl RustlsConnection {
pub async fn send_packet(&mut self, packet: &[u8]) -> io::Result<()> {
self.stream.write_all(packet).await?;
self.stream.flush().await?;
Ok(())
}
pub async fn send_raw(&mut self, data: &[u8]) -> io::Result<()> {
self.stream.write_all(data).await?;
self.stream.flush().await?;
Ok(())
}
pub async fn recv(&mut self, max_len: usize) -> io::Result<Vec<u8>> {
let mut buffer = vec![0u8; max_len];
let n = self.stream.read(&mut buffer).await?;
buffer.truncate(n);
Ok(buffer)
}
pub async fn close(mut self) -> io::Result<()> {
self.stream.shutdown().await
}
}
fn load_certs(reader: &mut dyn io::BufRead) -> io::Result<Vec<CertificateDer<'static>>> {
let mut pem_data = Vec::new();
reader.read_to_end(&mut pem_data)?;
if let Ok(certs) =
rustls_pemfile::certs(&mut pem_data.as_slice()).collect::<Result<Vec<_>, _>>()
{
if !certs.is_empty() {
return Ok(certs);
}
}
Ok(vec![CertificateDer::from(pem_data)])
}
fn load_private_key(reader: &mut dyn io::BufRead) -> io::Result<PrivateKeyDer<'static>> {
let mut key_data = Vec::new();
reader.read_to_end(&mut key_data)?;
if let Ok(Some(key)) = rustls_pemfile::private_key(&mut key_data.as_slice()) {
return Ok(key);
}
Ok(PrivateKeyDer::Pkcs8(key_data.into()))
}
#[derive(Debug)]
struct NoVerifier;
impl rustls::client::danger::ServerCertVerifier for NoVerifier {
fn verify_server_cert(
&self,
_end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn verify_tls13_signature(
&self,
_message: &[u8],
_cert: &CertificateDer<'_>,
_dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
Ok(rustls::client::danger::HandshakeSignatureValid::assertion())
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
vec![
rustls::SignatureScheme::RSA_PKCS1_SHA256,
rustls::SignatureScheme::RSA_PKCS1_SHA384,
rustls::SignatureScheme::RSA_PKCS1_SHA512,
rustls::SignatureScheme::ECDSA_NISTP256_SHA256,
rustls::SignatureScheme::ECDSA_NISTP384_SHA384,
rustls::SignatureScheme::ECDSA_NISTP521_SHA512,
rustls::SignatureScheme::RSA_PSS_SHA256,
rustls::SignatureScheme::RSA_PSS_SHA384,
rustls::SignatureScheme::RSA_PSS_SHA512,
rustls::SignatureScheme::ED25519,
]
}
}