#[cfg(any(feature = "agent", feature = "client"))]
use std::io;
use std::net::TcpStream;
#[cfg(feature = "client")]
const DEFAULT_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(60);
pub enum CradleStream<C> {
Plain(TcpStream),
Tls(Box<rustls::StreamOwned<C, TcpStream>>),
}
#[cfg(any(feature = "agent", feature = "client"))]
macro_rules! impl_cradle_stream_io {
($conn:ty) => {
impl io::Read for CradleStream<$conn> {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
match self {
Self::Plain(stream) => stream.read(buf),
Self::Tls(stream) => stream.read(buf),
}
}
}
impl io::Write for CradleStream<$conn> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
match self {
Self::Plain(stream) => stream.write(buf),
Self::Tls(stream) => stream.write(buf),
}
}
fn flush(&mut self) -> io::Result<()> {
match self {
Self::Plain(stream) => stream.flush(),
Self::Tls(stream) => stream.flush(),
}
}
}
impl CradleStream<$conn> {
pub fn set_read_timeout(&self, timeout: Option<std::time::Duration>) -> io::Result<()> {
match self {
Self::Plain(stream) => stream.set_read_timeout(timeout),
Self::Tls(stream) => stream.get_ref().set_read_timeout(timeout),
}
}
}
};
}
#[cfg(feature = "agent")]
impl_cradle_stream_io!(rustls::ServerConnection);
#[cfg(feature = "client")]
impl_cradle_stream_io!(rustls::ClientConnection);
#[cfg(feature = "agent")]
pub mod agent {
use super::install_crypto_provider;
use rustls::pki_types::{CertificateDer, PrivateKeyDer};
use rustls::ServerConfig;
use std::fs;
use std::io;
use std::net::TcpStream;
use std::path::Path;
use std::sync::Arc;
pub type AgentStream = super::CradleStream<rustls::ServerConnection>;
impl AgentStream {
pub fn plain(tcp: TcpStream) -> Self {
Self::Plain(tcp)
}
pub fn wrap_tls(tcp: TcpStream, config: &Arc<ServerConfig>) -> io::Result<Self> {
install_crypto_provider();
let conn = rustls::ServerConnection::new(config.clone()).map_err(io::Error::other)?;
Ok(Self::Tls(Box::new(rustls::StreamOwned::new(conn, tcp))))
}
}
pub fn load_server_config(cert_path: &Path, key_path: &Path) -> io::Result<Arc<ServerConfig>> {
install_crypto_provider();
let certs = load_certificates(cert_path)?;
let key = load_private_key(key_path)?;
let config = ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(certs, key)
.map_err(invalid_data)?;
Ok(Arc::new(config))
}
#[cfg(feature = "auto-cert")]
pub fn generate_self_signed() -> io::Result<Arc<ServerConfig>> {
install_crypto_provider();
let cert =
rcgen::generate_simple_self_signed(["cradle".to_string()]).map_err(io::Error::other)?;
let cert_der = CertificateDer::from(cert.cert.der().to_vec());
let key_der =
PrivateKeyDer::try_from(cert.signing_key.serialize_der()).map_err(io::Error::other)?;
let config = ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(vec![cert_der], key_der)
.map_err(io::Error::other)?;
Ok(Arc::new(config))
}
fn load_certificates(path: &Path) -> io::Result<Vec<CertificateDer<'static>>> {
let data = fs::read(path)?;
let certs = rustls_pemfile::certs(&mut data.as_slice())
.collect::<Result<Vec<_>, _>>()
.map_err(invalid_data)?;
if certs.is_empty() {
return Err(invalid_data("no certificates found in PEM file"));
}
Ok(certs)
}
fn load_private_key(path: &Path) -> io::Result<PrivateKeyDer<'static>> {
let data = fs::read(path)?;
rustls_pemfile::private_key(&mut data.as_slice())
.map_err(invalid_data)?
.ok_or_else(|| invalid_data("no private key found in PEM file"))
}
fn invalid_data(error: impl Into<Box<dyn std::error::Error + Send + Sync>>) -> io::Error {
io::Error::new(io::ErrorKind::InvalidData, error)
}
}
#[cfg(feature = "client")]
pub mod client {
use super::{install_crypto_provider, DEFAULT_READ_TIMEOUT};
use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
use rustls::{ClientConfig, DigitallySignedStruct, Error, SignatureScheme};
use std::io;
use std::net::TcpStream;
use std::sync::Arc;
pub type ClientStream = super::CradleStream<rustls::ClientConnection>;
pub fn connect(
addr: &str,
use_tls: bool,
insecure: bool,
pinned_fingerprint: Option<&str>,
) -> io::Result<ClientStream> {
install_crypto_provider();
let tcp = TcpStream::connect(addr)?;
tcp.set_read_timeout(Some(DEFAULT_READ_TIMEOUT))?;
if !use_tls {
return Ok(ClientStream::Plain(tcp));
}
let server_name = server_name_from_addr(addr)?;
let config = client_config(insecure, pinned_fingerprint)?;
let conn = rustls::ClientConnection::new(config, server_name).map_err(io::Error::other)?;
Ok(ClientStream::Tls(Box::new(rustls::StreamOwned::new(
conn, tcp,
))))
}
fn server_name_from_addr(addr: &str) -> io::Result<ServerName<'static>> {
let host = addr
.rsplit_once(':')
.map(|(host, _)| host)
.unwrap_or(addr)
.trim_matches(['[', ']']);
ServerName::try_from(host.to_string()).map_err(invalid_input)
}
fn client_config(
insecure: bool,
pinned_fingerprint: Option<&str>,
) -> io::Result<Arc<ClientConfig>> {
match (insecure, pinned_fingerprint) {
(true, _) => Ok(insecure_config()),
(false, Some(fingerprint)) => pinned_config(fingerprint),
(false, None) => Ok(default_config()),
}
}
fn default_config() -> Arc<ClientConfig> {
let mut root_store = rustls::RootCertStore::empty();
root_store.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
Arc::new(
ClientConfig::builder()
.with_root_certificates(root_store)
.with_no_client_auth(),
)
}
fn insecure_config() -> Arc<ClientConfig> {
Arc::new(
ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(AcceptAnyCert))
.with_no_client_auth(),
)
}
fn pinned_config(fingerprint: &str) -> io::Result<Arc<ClientConfig>> {
let expected = parse_sha256_fingerprint(fingerprint)?;
Ok(Arc::new(
ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(PinnedCertVerifier { expected }))
.with_no_client_auth(),
))
}
#[derive(Debug)]
struct AcceptAnyCert;
impl ServerCertVerifier for AcceptAnyCert {
fn verify_server_cert(
&self,
_end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp: &[u8],
_now: UnixTime,
) -> Result<ServerCertVerified, Error> {
Ok(ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
verify_tls12_signature(message, cert, dss)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
verify_tls13_signature(message, cert, dss)
}
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
supported_verify_schemes()
}
}
#[derive(Debug)]
struct PinnedCertVerifier {
expected: [u8; 32],
}
impl ServerCertVerifier for PinnedCertVerifier {
fn verify_server_cert(
&self,
end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp: &[u8],
_now: UnixTime,
) -> Result<ServerCertVerified, Error> {
let actual = ring::digest::digest(&ring::digest::SHA256, end_entity.as_ref());
if actual.as_ref() == self.expected {
return Ok(ServerCertVerified::assertion());
}
Err(Error::General("certificate fingerprint mismatch".into()))
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
verify_tls12_signature(message, cert, dss)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
verify_tls13_signature(message, cert, dss)
}
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
supported_verify_schemes()
}
}
fn verify_tls12_signature(
message: &[u8],
cert: &CertificateDer,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn verify_tls13_signature(
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn supported_verify_schemes() -> Vec<SignatureScheme> {
rustls::crypto::ring::default_provider()
.signature_verification_algorithms
.supported_schemes()
}
fn parse_sha256_fingerprint(fingerprint: &str) -> io::Result<[u8; 32]> {
let hex = fingerprint.replace(':', "");
if hex.len() != 64 {
return Err(invalid_input(format!(
"fingerprint must be 32 SHA-256 bytes, got {}",
hex.len() / 2
)));
}
let mut bytes = [0u8; 32];
for (idx, chunk) in hex.as_bytes().chunks_exact(2).enumerate() {
let chunk = std::str::from_utf8(chunk).map_err(invalid_input)?;
bytes[idx] = u8::from_str_radix(chunk, 16).map_err(invalid_input)?;
}
Ok(bytes)
}
fn invalid_input(error: impl Into<Box<dyn std::error::Error + Send + Sync>>) -> io::Error {
io::Error::new(io::ErrorKind::InvalidInput, error)
}
}
#[cfg(any(feature = "agent", feature = "client"))]
fn install_crypto_provider() {
let _ = rustls::crypto::ring::default_provider().install_default();
}