cradle-shared 0.1.0

Shared utilities in use by the Cradle agent and cli
Documentation
#[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);

/// Generic stream object for the cradle workspace
///
/// Most users should use [`agent::AgentStream`] or [`client::ClientStream`]
pub enum CradleStream<C> {
    /// Specifies a simple TCP stream
    Plain(TcpStream),
    /// Specifies a TLS stream where the generic, `C`, is the client/server settings
    Tls(Box<rustls::StreamOwned<C, TcpStream>>),
}

/// Macro that implements some of the `std::io` functionality for the streams
#[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> {
            /// Optionally sets the read timeout duration
            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")]
/// Contains the TLS logic for the cradle 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;
    /// An Agent connection stream
    pub type AgentStream = super::CradleStream<rustls::ServerConnection>;
    impl AgentStream {
        /// Helper to create an `AgentStream::Plain` object
        pub fn plain(tcp: TcpStream) -> Self {
            Self::Plain(tcp)
        }

        /// Wraps the given `TcpStream` with a TLS envelope as specified by the `config` variable
        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))))
        }
    }

    /// Loads the server TLS parameters: Certificate and Key paths
    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")]
    /// Auto-generates the certificates to use
    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")]
/// Contains the TLS logic for the cradle 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;

    /// A Client connection stream
    pub type ClientStream = super::CradleStream<rustls::ClientConnection>;

    /// Connects to the specified address.
    /// Allows specifying whether to use TLS, allow insecure connections, and a pinned fingerprint for extra security
    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();
}