diff --git a/Cargo.lock b/Cargo.lock index 1e740c5..d5c0827 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -108,12 +108,6 @@ dependencies = [ "windows-targets", ] -[[package]] -name = "base64" -version = "0.21.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9d297deb1925b89f2ccc13d7635fa0714f12c87adce1c75356b39ca9b7178567" - [[package]] name = "base64" version = "0.22.1" @@ -1005,7 +999,7 @@ version = "3.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "38af38e8470ac9dee3ce1bae1af9c1671fffc44ddfd8bd1d0a3445bf349a8ef3" dependencies = [ - "base64 0.22.1", + "base64", "serde", ] @@ -1351,35 +1345,38 @@ dependencies = [ [[package]] name = "rustls" -version = "0.21.12" +version = "0.22.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3f56a14d1f48b391359b22f731fd4bd7e43c97f3c50eee276f3aa09c94784d3e" +checksum = "bf4ef73721ac7bcd79b2b315da7779d8fc09718c6b3d2d1b2d94850eb8c18432" dependencies = [ "log", "ring", + "rustls-pki-types", "rustls-webpki", - "sct", + "subtle", + "zeroize", ] [[package]] name = "rustls-native-certs" -version = "0.6.3" +version = "0.7.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a9aace74cb666635c918e9c12bc0d348266037aa8eb599b5cba565709a8dff00" +checksum = "e5bfb394eeed242e909609f56089eecfe5fda225042e8b171791b9c95f5931e5" dependencies = [ "openssl-probe", "rustls-pemfile", + "rustls-pki-types", "schannel", "security-framework", ] [[package]] name = "rustls-pemfile" -version = "1.0.4" +version = "2.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1c74cae0a4cf6ccbbf5f359f08efdf8ee7e1dc532573bf0db71968cb56b1448c" +checksum = "dce314e5fee3f39953d46bb63bb8a46d40c2f8fb7cc5a3b6cab2bde9721d6e50" dependencies = [ - "base64 0.21.7", + "rustls-pki-types", ] [[package]] @@ -1393,11 +1390,12 @@ dependencies = [ [[package]] name = "rustls-webpki" -version = "0.101.7" +version = "0.102.8" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8b6275d1ee7a1cd780b64aca7726599a1dbc893b1e64144529e55c3c2f745765" +checksum = "64ca1bc8749bd4cf37b5ce386cc146580777b4e8572c7b97baf22c83f444bee9" dependencies = [ "ring", + "rustls-pki-types", "untrusted", ] @@ -1449,16 +1447,6 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" -[[package]] -name = "sct" -version = "0.7.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "da046153aa2352493d6cb7da4b6e5c0c057d8a1d0a9aa8560baffdd945acd414" -dependencies = [ - "ring", - "untrusted", -] - [[package]] name = "security-framework" version = "2.11.1" @@ -1784,11 +1772,12 @@ dependencies = [ [[package]] name = "tokio-rustls" -version = "0.24.1" +version = "0.25.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c28327cf380ac148141087fbfb9de9d7bd4e84ab5d2c28fbc911d753de8a7081" +checksum = "775e0c0f0adb3a2f22a00c4745d728b479985fc15ee7ca6a2608388c5569860f" dependencies = [ "rustls", + "rustls-pki-types", "tokio", ] diff --git a/Cargo.toml b/Cargo.toml index 98a9144..ad2f34f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -90,11 +90,11 @@ zstd = "0.13" toml = "0.8" # TLS dependencies -rustls = { version = "0.21", features = ["dangerous_configuration"] } -tokio-rustls = "0.24" +rustls = { version = "0.22" } +tokio-rustls = "0.25" rcgen = { version = "0.14.5", features = ["pem"] } -rustls-pemfile = "1.0" -rustls-native-certs = "0.6" +rustls-pemfile = "2.1" +rustls-native-certs = "0.7" time = "0.3.47" [dev-dependencies] diff --git a/src/service/tls_client.rs b/src/service/tls_client.rs index 691eb52..d927526 100644 --- a/src/service/tls_client.rs +++ b/src/service/tls_client.rs @@ -1,4 +1,5 @@ use futures::{SinkExt, StreamExt}; +use rustls::pki_types::ServerName; use std::sync::Arc; use tokio::net::TcpStream; use tokio_rustls::client::TlsStream; @@ -7,7 +8,7 @@ use tracing::{debug, instrument}; use crate::core::codec::PacketCodec; use crate::core::packet::Packet; -use crate::error::Result; +use crate::error::{ProtocolError, Result}; use crate::protocol::message::Message; use crate::transport::session_cache::SessionCache; use crate::transport::tls::TlsClientConfig; @@ -48,7 +49,6 @@ impl TlsClient { /// Some(Arc::new(cache)) /// ).await?; /// ``` - #[instrument(skip(config, session_cache))] pub async fn connect_with_session( addr: &str, config: TlsClientConfig, @@ -58,7 +58,13 @@ impl TlsClient { let connector = tokio_rustls::TlsConnector::from(std::sync::Arc::new(tls_config)); let stream = TcpStream::connect(addr).await?; - let domain = config.server_name()?; + + // Create ServerName from owned string to ensure 'static lifetime + // Note: Box::leak() is used here to satisfy tokio_rustls' 'static requirement + let server_name_str = config.server_name_string(); + let domain_static: &'static str = Box::leak(server_name_str.into_boxed_str()); + let domain = ServerName::try_from(domain_static) + .map_err(|_| ProtocolError::TlsError("Invalid server name".into()))?; let tls_stream = connector.connect(domain, stream).await?; let framed = Framed::new(tls_stream, PacketCodec); diff --git a/src/transport/local.rs b/src/transport/local.rs index c8e026e..6b40e12 100644 --- a/src/transport/local.rs +++ b/src/transport/local.rs @@ -317,10 +317,7 @@ fn convert_to_pipe_name(path: &str) -> String { } // Extract a meaningful name from the path - let name = path - .trim_start_matches('/') - .replace('/', "_") - .replace('\\', "_"); + let name = path.trim_start_matches('/').replace(['/', '\\'], "_"); // Use a default if empty let name = if name.is_empty() { diff --git a/src/transport/tls.rs b/src/transport/tls.rs index 5a433f2..7941025 100644 --- a/src/transport/tls.rs +++ b/src/transport/tls.rs @@ -21,8 +21,9 @@ use std::net::SocketAddr; use std::path::Path; use std::sync::Arc; -use rustls::ServerName; -use rustls::{Certificate, ClientConfig, PrivateKey, RootCertStore, ServerConfig}; +use rustls::client::danger::{ServerCertVerified, ServerCertVerifier}; +use rustls::pki_types::{CertificateDer, PrivateKeyDer, ServerName, UnixTime}; +use rustls::{ClientConfig, DigitallySignedStruct, RootCertStore, ServerConfig}; use rustls_pemfile::{certs, pkcs8_private_keys}; use tokio::net::{TcpListener, TcpStream}; use tokio_rustls::client::TlsStream as ClientTlsStream; @@ -37,52 +38,129 @@ use crate::error::{ProtocolError, Result}; use futures::{SinkExt, StreamExt}; // Custom certificate verifiers +#[derive(Debug)] struct CertificateFingerprint { fingerprint: Vec, } -impl rustls::client::ServerCertVerifier for CertificateFingerprint { +impl ServerCertVerifier for CertificateFingerprint { fn verify_server_cert( &self, - end_entity: &Certificate, - _intermediates: &[Certificate], + end_entity: &CertificateDer<'_>, + _intermediates: &[CertificateDer<'_>], _server_name: &ServerName, - _scts: &mut dyn Iterator, _ocsp_response: &[u8], - _now: std::time::SystemTime, - ) -> std::result::Result { + _now: UnixTime, + ) -> std::result::Result { use sha2::{Digest, Sha256}; let mut hasher = Sha256::new(); - hasher.update(&end_entity.0); + hasher.update(end_entity); let hash = hasher.finalize(); if hash.as_slice() == self.fingerprint.as_slice() { - Ok(rustls::client::ServerCertVerified::assertion()) + Ok(ServerCertVerified::assertion()) } else { Err(rustls::Error::General( "Pinned certificate hash mismatch".into(), )) } } + + fn verify_tls12_signature( + &self, + _message: &[u8], + _cert: &CertificateDer<'_>, + _dss: &DigitallySignedStruct, + ) -> std::result::Result { + Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) + } + + fn verify_tls13_signature( + &self, + _message: &[u8], + _cert: &CertificateDer<'_>, + _dss: &DigitallySignedStruct, + ) -> std::result::Result { + Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) + } + + fn supported_verify_schemes(&self) -> Vec { + // Accept common signature schemes + vec![ + rustls::SignatureScheme::RSA_PKCS1_SHA256, + rustls::SignatureScheme::ECDSA_NISTP256_SHA256, + rustls::SignatureScheme::ED25519, + ] + } } +#[derive(Debug)] struct AcceptAnyServerCert; -impl rustls::client::ServerCertVerifier for AcceptAnyServerCert { +impl ServerCertVerifier for AcceptAnyServerCert { fn verify_server_cert( &self, - _end_entity: &Certificate, - _intermediates: &[Certificate], + _end_entity: &CertificateDer<'_>, + _intermediates: &[CertificateDer<'_>], _server_name: &ServerName, - _scts: &mut dyn Iterator, _ocsp_response: &[u8], - _now: std::time::SystemTime, - ) -> std::result::Result { - Ok(rustls::client::ServerCertVerified::assertion()) + _now: UnixTime, + ) -> std::result::Result { + Ok(ServerCertVerified::assertion()) + } + + fn verify_tls12_signature( + &self, + _message: &[u8], + _cert: &CertificateDer<'_>, + _dss: &DigitallySignedStruct, + ) -> std::result::Result { + Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) + } + + fn verify_tls13_signature( + &self, + _message: &[u8], + _cert: &CertificateDer<'_>, + _dss: &DigitallySignedStruct, + ) -> std::result::Result { + Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) + } + + fn supported_verify_schemes(&self) -> Vec { + vec![ + rustls::SignatureScheme::RSA_PKCS1_SHA256, + rustls::SignatureScheme::ECDSA_NISTP256_SHA256, + rustls::SignatureScheme::ED25519, + ] } } +/// Helper function to load a private key from PKCS8 format +fn load_private_key(reader: &mut BufReader) -> Result> { + // Try to load PKCS8 keys + // Seek to beginning of file first + reader + .seek(std::io::SeekFrom::Start(0)) + .map_err(ProtocolError::Io)?; + + // pkcs8_private_keys returns an iterator of Results + let keys: std::result::Result, _> = pkcs8_private_keys(reader).collect(); + let keys = + keys.map_err(|_| ProtocolError::TlsError("Failed to parse PKCS8 private key".into()))?; + + if !keys.is_empty() { + return Ok(PrivateKeyDer::Pkcs8(keys[0].clone_key())); + } + + // Note: Add support for other key formats like RSA or EC if needed + + Err(ProtocolError::TlsError( + "No supported private key format found".into(), + )) +} + /// TLS protocol version pub enum TlsVersion { /// TLS 1.2 @@ -185,28 +263,22 @@ impl TlsServerConfig { let cert_file = File::open(&self.cert_path) .map_err(|e| ProtocolError::TlsError(format!("Failed to open cert file: {e}")))?; let mut cert_reader = BufReader::new(cert_file); - let cert_chain = certs(&mut cert_reader) + let cert_chain: std::result::Result, _> = certs(&mut cert_reader).collect(); + let cert_chain: Vec> = cert_chain .map_err(|_| ProtocolError::TlsError("Failed to parse certificate".into()))?; - // Convert to rustls Certificate type - let cert_chain: Vec = cert_chain.into_iter().map(Certificate).collect(); + if cert_chain.is_empty() { + return Err(ProtocolError::TlsError("No certificates found".into())); + } // Load private key let key_file = File::open(&self.key_path) .map_err(|e| ProtocolError::TlsError(format!("Failed to open key file: {e}")))?; let mut key_reader = BufReader::new(key_file); - let keys = pkcs8_private_keys(&mut key_reader) - .map_err(|_| ProtocolError::TlsError("Failed to parse private key".into()))?; - - if keys.is_empty() { - return Err(ProtocolError::TlsError("No private keys found".into())); - } - - // Convert to rustls PrivateKey - let private_key = PrivateKey(keys[0].clone()); + let private_key = load_private_key(&mut key_reader)?; // Validate TLS versions if specified - // Note: In rustls 0.21, with_safe_defaults() restricts to TLS 1.2+ (best practice) + // Note: In rustls 0.22, with_safe_defaults() restricts to TLS 1.2+ (best practice) if let Some(versions) = &self.tls_versions { let mut has_tls13 = false; let mut has_tls12 = false; @@ -228,19 +300,17 @@ impl TlsServerConfig { } // Create a server configuration with safe defaults (TLS 1.2+, modern ciphersuites) - let config_builder = ServerConfig::builder().with_safe_defaults(); - - // Note: rustls 0.21 doesn't expose API to change ciphersuites after builder creation. - // with_safe_defaults() already provides excellent defaults: - // - TLS 1.3: CHACHA20_POLY1305, AES_128_GCM, AES_256_GCM - // - TLS 1.2: same + some legacy suites - // Custom cipher suite restriction requires building at compile time. + let config_builder = ServerConfig::builder_with_provider(std::sync::Arc::new( + rustls::crypto::ring::default_provider(), + )) + .with_safe_default_protocol_versions() + .map_err(|_| ProtocolError::TlsError("Failed to configure TLS protocol versions".into()))?; let cert_builder = config_builder.with_no_client_auth(); // Build config with certificates let mut config = cert_builder - .with_single_cert(cert_chain.clone(), private_key.clone()) + .with_single_cert(cert_chain.clone(), private_key.clone_key()) .map_err(|e| ProtocolError::TlsError(format!("TLS error: {e}")))?; // Configure client authentication if required (mTLS) @@ -250,34 +320,48 @@ impl TlsServerConfig { ProtocolError::TlsError(format!("Failed to open client CA file: {e}")) })?; let mut client_ca_reader = BufReader::new(client_ca_file); - let client_ca_certs = certs(&mut client_ca_reader).map_err(|_| { + let client_ca_certs: std::result::Result, _> = + certs(&mut client_ca_reader).collect(); + let client_ca_certs: Vec> = client_ca_certs.map_err(|_| { ProtocolError::TlsError("Failed to parse client CA certificate".into()) })?; - // Convert to rustls Certificate type - let client_ca_certs: Vec = - client_ca_certs.into_iter().map(Certificate).collect(); + if client_ca_certs.is_empty() { + return Err(ProtocolError::TlsError( + "No client CA certificates found".into(), + )); + } // Create client cert verifier let mut client_root_store = RootCertStore::empty(); - for cert in &client_ca_certs { + for cert in client_ca_certs { client_root_store.add(cert).map_err(|e| { ProtocolError::TlsError(format!("Failed to add client CA cert: {e}")) })?; } - // Create client authentication verifier - let client_auth = Arc::new(rustls::server::AllowAnyAuthenticatedClient::new( + // Create client authentication verifier using WebPkiClientVerifier + let client_auth = rustls::server::WebPkiClientVerifier::builder(std::sync::Arc::new( client_root_store, - )); + )) + .build() + .map_err(|e| { + ProtocolError::TlsError(format!("Failed to build client verifier: {e}")) + })?; // Create new config builder with client auth - let new_builder = ServerConfig::builder().with_safe_defaults(); + let new_builder = ServerConfig::builder_with_provider(std::sync::Arc::new( + rustls::crypto::ring::default_provider(), + )) + .with_safe_default_protocol_versions() + .map_err(|_| { + ProtocolError::TlsError("Failed to configure TLS protocol versions".into()) + })?; let new_cert_builder = new_builder.with_client_cert_verifier(client_auth); // Build a new config with certificates and client auth config = new_cert_builder - .with_single_cert(cert_chain, private_key) + .with_single_cert(cert_chain, private_key.clone_key()) .map_err(|e| ProtocolError::TlsError(format!("TLS error with client auth: {e}")))?; debug!("mTLS enabled with client certificate verification required"); @@ -294,6 +378,14 @@ impl TlsServerConfig { Ok(config) } + + /// Calculate SHA-256 hash for a certificate to use with pinning + pub fn calculate_cert_hash(cert: &CertificateDer<'_>) -> Vec { + use sha2::{Digest, Sha256}; + let mut hasher = Sha256::new(); + hasher.update(cert.as_ref()); + hasher.finalize().to_vec() + } } /// TLS Client Configuration @@ -380,37 +472,6 @@ impl TlsClientConfig { self } - /// Calculate SHA-256 hash for a certificate to use with pinning - pub fn calculate_cert_hash(cert: &Certificate) -> Vec { - use sha2::{Digest, Sha256}; - let mut hasher = Sha256::new(); - hasher.update(&cert.0); - hasher.finalize().to_vec() - } - - /// Helper method to load a private key from PKCS8 format - fn load_private_key(reader: &mut BufReader) -> Result { - // Try to load PKCS8 keys - // Seek to beginning of file first - reader - .seek(std::io::SeekFrom::Start(0)) - .map_err(ProtocolError::Io)?; - - // We need to use pkcs8_private_keys on the BufReader directly since it implements BufRead - let keys = pkcs8_private_keys(reader) - .map_err(|_| ProtocolError::TlsError("Failed to parse PKCS8 private key".into()))?; - - if !keys.is_empty() { - return Ok(PrivateKey(keys[0].clone())); - } - - // Note: Add support for other key formats like RSA or EC if needed - - Err(ProtocolError::TlsError( - "No supported private key format found".into(), - )) - } - /// Load the TLS client configuration pub fn load_client_config(&self) -> Result { self.log_tls_version_info(); @@ -447,9 +508,12 @@ impl TlsClientConfig { /// Build secure client config with system root CAs fn build_secure_client_config(&self) -> Result { let root_store = self.load_system_root_certificates()?; - let builder = ClientConfig::builder() - .with_safe_defaults() - .with_root_certificates(root_store); + let builder = ClientConfig::builder_with_provider(std::sync::Arc::new( + rustls::crypto::ring::default_provider(), + )) + .with_safe_default_protocol_versions() + .map_err(|_| ProtocolError::TlsError("Failed to configure TLS protocol versions".into()))? + .with_root_certificates(root_store); // Apply client auth directly if let (Some(client_cert_path), Some(client_key_path)) = @@ -467,9 +531,15 @@ impl TlsClientConfig { /// Build insecure client config with custom verifier fn build_insecure_client_config(&self) -> Result { - let builder = ClientConfig::builder().with_safe_defaults(); + let builder = ClientConfig::builder_with_provider(std::sync::Arc::new( + rustls::crypto::ring::default_provider(), + )) + .with_safe_default_protocol_versions() + .map_err(|_| ProtocolError::TlsError("Failed to configure TLS protocol versions".into()))?; let verifier = self.create_custom_verifier(); - let custom_builder = builder.with_custom_certificate_verifier(verifier); + let custom_builder = builder + .dangerous() + .with_custom_certificate_verifier(verifier); // Apply client auth directly if let (Some(client_cert_path), Some(client_key_path)) = @@ -494,7 +564,7 @@ impl TlsClientConfig { .map_err(|e| ProtocolError::TlsError(format!("Failed to load native certs: {e}")))?; for cert in native_certs { - root_store.add(&Certificate(cert.0)).map_err(|e| { + root_store.add(cert).map_err(|e| { ProtocolError::TlsError(format!("Failed to add cert to root store: {e}")) })?; } @@ -503,7 +573,7 @@ impl TlsClientConfig { } /// Create custom certificate verifier (pinning or accept-any) - fn create_custom_verifier(&self) -> Arc { + fn create_custom_verifier(&self) -> Arc { if let Some(hash) = &self.pinned_cert_hash { Arc::new(CertificateFingerprint { fingerprint: hash.clone(), @@ -518,11 +588,13 @@ impl TlsClientConfig { &self, cert_path: &str, key_path: &str, - ) -> Result<(Vec, PrivateKey)> { + ) -> Result<(Vec>, PrivateKeyDer<'static>)> { // Load certificate let cert_file = File::open(cert_path).map_err(ProtocolError::Io)?; let mut cert_reader = BufReader::new(cert_file); - let certs = rustls_pemfile::certs(&mut cert_reader) + let certs_result: std::result::Result, _> = + rustls_pemfile::certs(&mut cert_reader).collect(); + let certs: Vec> = certs_result .map_err(|_| ProtocolError::TlsError("Failed to parse client certificate".into()))?; if certs.is_empty() { @@ -534,17 +606,21 @@ impl TlsClientConfig { // Load private key let key_file = File::open(key_path).map_err(ProtocolError::Io)?; let mut key_reader = BufReader::new(key_file); - let key = Self::load_private_key(&mut key_reader)?; + let key = load_private_key(&mut key_reader)?; - let cert_chain = certs.into_iter().map(Certificate).collect(); - Ok((cert_chain, key)) + Ok((certs, key)) } /// Get the server name as a rustls::ServerName - pub fn server_name(&self) -> Result { + pub fn server_name(&self) -> Result> { ServerName::try_from(self.server_name.as_str()) .map_err(|_| ProtocolError::TlsError("Invalid server name".into())) } + + /// Get the server name as an owned String + pub fn server_name_string(&self) -> String { + self.server_name.clone() + } } /// Start a TLS server on the given address @@ -619,7 +695,6 @@ where } /// Connect to a TLS server -#[instrument(skip(config), fields(address=%addr))] pub async fn connect( addr: &str, config: TlsClientConfig, @@ -628,7 +703,13 @@ pub async fn connect( let connector = TlsConnector::from(tls_config); let stream = TcpStream::connect(addr).await?; - let domain = config.server_name()?; + + // Create ServerName from owned string to ensure 'static lifetime + // Note: Box::leak() is used here to satisfy tokio_rustls' 'static requirement + let server_name_str = config.server_name_string(); + let domain_static: &'static str = Box::leak(server_name_str.into_boxed_str()); + let domain = ServerName::try_from(domain_static) + .map_err(|_| ProtocolError::TlsError("Invalid server name".into()))?; let tls_stream = connector .connect(domain, stream)