diff --git a/Cargo.lock b/Cargo.lock index f06f474..b9e3477 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -686,15 +686,11 @@ dependencies = [ "hyper", "hyper-util", "jsonrpsee", - "p256", "pccs", "pem-rfc7468", "pin-project-lite", - "pkcs1", - "pkcs8", "rcgen", "reqwest 0.13.4", - "rsa", "rustls-pemfile", "serde", "serde_json", @@ -707,7 +703,27 @@ dependencies = [ "tracing", "tracing-subscriber", "webpki-roots 1.0.4", - "x509-parser", +] + +[[package]] +name = "attested-tls-tcp-tunnel" +version = "0.1.0" +dependencies = [ + "anyhow", + "attested-tls", + "bytes", + "clap", + "h2", + "http", + "rustls-pemfile", + "serde_json", + "tempfile", + "thiserror 2.0.17", + "tokio", + "tokio-rustls", + "tracing", + "tracing-subscriber", + "webpki-roots 1.0.4", ] [[package]] @@ -5169,7 +5185,6 @@ dependencies = [ "bytes", "libc", "mio", - "parking_lot", "pin-project-lite", "signal-hook-registry", "socket2 0.6.4", diff --git a/Cargo.toml b/Cargo.toml index 95317a5..e53e81e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,5 +1,27 @@ [workspace] -members = [".", "attested-tls"] +members = [".", "attested-tls", "tcp-tunnel"] + +[workspace.dependencies] +anyhow = "1.0.100" +attestation = { git = "https://github.com/flashbots/attested-tls", branch = "main" } +attested-tls = { path = "attested-tls", default-features = false } +bytes = "1.11.1" +clap = { version = "4.5.51", features = ["derive", "env"] } +h2 = "0.4.12" +http = "1.3.1" +http-body-util = "0.1.3" +hyper = { version = "1.7.0", features = ["http2"] } +hyper-util = { version = "0.1.17", features = ["tokio"] } +rcgen = "0.14.5" +rustls-pemfile = "2.2.0" +serde_json = "1.0.145" +tempfile = "3.23.0" +thiserror = "2.0.17" +tokio = "1.48.0" +tokio-rustls = { version = "0.26.4", default-features = false } +tracing = "0.1.41" +tracing-subscriber = { version = "0.3.20", features = ["env-filter", "json"] } +webpki-roots = "1.0.4" [package] name = "attested-tls-proxy" @@ -11,45 +33,39 @@ repository = "https://github.com/flashbots/attested-tls-proxy" keywords = ["attested-TLS", "CVM", "TDX"] [dependencies] -attested-tls = { path = "attested-tls", default-features = false } -tokio = { version = "1.48.0", features = ["full"] } -tokio-rustls = { version = "0.26.4", default-features = false, features = [ - "aws_lc_rs", -] } -x509-parser = { version = "0.18.0", features = ["verify"] } -thiserror = "2.0.17" -clap = { version = "4.5.51", features = ["derive", "env"] } -rustls-pemfile = "2.2.0" -anyhow = "1.0.100" +attested-tls = { workspace = true, features = ["self-signed"] } +tokio = { workspace = true, features = ["fs", "io-std", "io-util", "macros", "net", "rt-multi-thread", "sync", "time"] } +tokio-rustls = { workspace = true, features = ["aws_lc_rs"] } +thiserror.workspace = true +clap.workspace = true +rustls-pemfile.workspace = true +anyhow.workspace = true pem-rfc7468 = { version = "0.7.0", features = ["std"] } -hyper = { version = "1.7.0", features = ["server", "http2"] } -h2 = "0.4.12" -hyper-util = { version = "0.1.17", features = ["tokio"] } -http-body-util = "0.1.3" -bytes = "1.11.1" -http = "1.3.1" -serde_json = "1.0.145" +hyper = { workspace = true, features = ["server"] } +h2.workspace = true +hyper-util.workspace = true +http-body-util.workspace = true +bytes.workspace = true +http.workspace = true +serde_json.workspace = true serde = "1.0.228" reqwest = { version = "0.13.4", default-features = false, features = [ "rustls-no-provider", ] } -webpki-roots = "1.0.4" -tracing = "0.1.41" -tracing-subscriber = { version = "0.3.20", features = ["env-filter", "json"] } +webpki-roots.workspace = true +tracing.workspace = true +tracing-subscriber.workspace = true axum = "0.8.8" tower-http = { version = "0.6.7", features = ["fs"] } -rsa = { version = "0.9", default-features = false } -p256 = { version = "0.13.2", features = ["pkcs8"] } -pkcs1 = "0.7.5" -pkcs8 = "0.10.2" -rcgen = "0.14.5" pin-project-lite = "0.2.16" pccs = { git = "https://github.com/flashbots/attested-tls", branch = "main" } [dev-dependencies] -tempfile = "3.23.0" +tokio = { workspace = true, features = ["process"] } +rcgen.workspace = true +tempfile.workspace = true tdx-quote = { version = "0.0.5", features = ["mock"] } -attested-tls = { path = "attested-tls", features = ["test-helpers", "mock"] } +attested-tls = { workspace = true, default-features = true, features = ["test-helpers", "mock"] } jsonrpsee = { version = "0.26.0", features = ["server"] } [features] diff --git a/README.md b/README.md index 0d03647..605f10e 100644 --- a/README.md +++ b/README.md @@ -17,6 +17,8 @@ It has five subcommands: - `attested-tls-proxy attested-file-server` - serve files from a local filesystem path over an attested TLS channel. - `attested-tls-proxy attested-get` - connect to a proxy server, verify its attestation, make a single HTTP GET request, and write the response body to standard output. +If you rather want opaque TCP forwarding, see [`attested-tls-tcp-tunnel`](tcp-tunnel/README.md). It provides a separate CLI and library with one attested connection per source TCP connection. + ### How it works This works as follows: diff --git a/attested-tls/Cargo.toml b/attested-tls/Cargo.toml index ef4930d..07cd0d3 100644 --- a/attested-tls/Cargo.toml +++ b/attested-tls/Cargo.toml @@ -8,17 +8,17 @@ repository = "https://github.com/flashbots/attested-tls-proxy" keywords = ["attested-TLS", "CVM", "TDX"] [dependencies] -tokio = { version = "1.48.0", features = ["full"] } -tokio-rustls = { version = "0.26.4", default-features = false } +tokio = { workspace = true, features = ["io-util", "net", "rt"] } +tokio-rustls.workspace = true sha2 = "0.10.9" x509-parser = "0.18.0" -thiserror = "2.0.17" -webpki-roots = "1.0.4" -http = "1.3.1" -serde_json = "1.0.145" -tracing = "0.1.41" +thiserror.workspace = true +webpki-roots.workspace = true +http.workspace = true +serde_json.workspace = true +tracing.workspace = true parity-scale-codec = "3.7.5" -attestation = { git = "https://github.com/flashbots/attested-tls", branch = "main" } +attestation.workspace = true # Used for websocket support tokio-tungstenite = { version = "0.28.0", optional = true } @@ -29,22 +29,26 @@ alloy-rpc-client = { version = "1.1.3", optional = true } tower-service = { version = "0.3.3", optional = true } alloy-transport-http = { version = "1.4.3", features = ["hyper"], optional = true } url = { version = "2.5.7", optional = true } -hyper = { version = "1.7.0", features = ["client", "http2"], optional = true } -hyper-util = { version = "0.1.17", features = ["tokio"], optional = true } -bytes = { version = "1.11.1", optional = true } -http-body-util = { version = "0.1.3", optional = true } +hyper = { workspace = true, optional = true, features = ["client"] } +hyper-util = { workspace = true, optional = true } +bytes = { workspace = true, optional = true } +http-body-util = { workspace = true, optional = true } -# Used by test helpers -rcgen = { version = "0.14.5", optional = true } +# Used by test helpers and self-signed certificate support +rcgen = { workspace = true, optional = true } [dev-dependencies] -rcgen = "0.14.5" -tempfile = "3.23.0" -attestation = { git = "https://github.com/flashbots/attested-tls", branch = "main", features = ["mock"] } +tokio = { workspace = true, features = ["macros", "time"] } +rcgen.workspace = true +tempfile.workspace = true +attestation = { workspace = true, features = ["mock"] } [features] default = ["ws", "rpc"] +# Self-signed certificate generation and verification. +self-signed = ["rcgen", "x509-parser/verify"] + # Adds support for Microsoft Azure attestation generation and verification azure = ["attestation/azure"] @@ -53,6 +57,7 @@ ws = ["tokio-tungstenite", "futures-util"] # Adds JSON RPC support rpc = [ + "tokio/sync", "alloy-rpc-client", "tower-service", "alloy-transport-http", diff --git a/attested-tls/README.md b/attested-tls/README.md index d699e5b..71229e2 100644 --- a/attested-tls/README.md +++ b/attested-tls/README.md @@ -14,6 +14,8 @@ It uses session binding through exported key material from the TLS session. This Attestation may be provided by either the server, or the client, or both. +The optional `self-signed` feature exposes the `self_signed` module for generating and verifying self-signed certificates. + ## Protocol Specification A TLS 1.3 handshake is made between server and client. The protocol name `flashbots-ratls/1` is included in ALPN. Future versions of the protocol may add additional protocol names which increment the number given after the slash, but backwards compatibility will be provided through also specifying `flashbots-ratls/1`. diff --git a/attested-tls/src/lib.rs b/attested-tls/src/lib.rs index 7cc68c7..12f7b31 100644 --- a/attested-tls/src/lib.rs +++ b/attested-tls/src/lib.rs @@ -8,6 +8,9 @@ pub mod attested_rpc; #[cfg(any(test, feature = "test-helpers"))] pub mod test_helpers; +#[cfg(feature = "self-signed")] +pub mod self_signed; + pub use attestation; use attestation::{ diff --git a/attested-tls/src/self_signed.rs b/attested-tls/src/self_signed.rs new file mode 100644 index 0000000..264ede2 --- /dev/null +++ b/attested-tls/src/self_signed.rs @@ -0,0 +1,324 @@ +use std::{net::IpAddr, sync::Arc}; +use tokio_rustls::rustls::{ + self, + crypto::CryptoProvider, + pki_types::{self, CertificateDer, PrivatePkcs8KeyDer}, +}; +use x509_parser::prelude::{FromDer, X509Certificate}; + +use crate::{AttestedTlsError, TlsCertAndKey}; + +/// Generate a self signed certifcate +pub fn generate_self_signed_cert(ip_address: IpAddr) -> Result { + let keypair = rcgen::KeyPair::generate()?; + let mut params = rcgen::CertificateParams::default(); + params + .subject_alt_names + .push(rcgen::SanType::IpAddress(ip_address)); + + let cert = params.self_signed(&keypair)?; + Ok(TlsCertAndKey { + cert_chain: vec![cert.der().clone()], + key: PrivatePkcs8KeyDer::from(keypair.serialize_der()).into(), + }) +} + +/// Client TLS configuration which accepts self-signed remote certificates +pub fn client_tls_config_allow_self_signed() -> Result { + Ok(rustls::ClientConfig::builder() + .dangerous() + .with_custom_certificate_verifier(SkipServerVerification::new()?) + .with_no_client_auth()) +} + +/// Used to allow verification of self-signed certificates +#[derive(Debug, Clone)] +pub struct SkipServerVerification { + supported_algs: rustls::crypto::WebPkiSupportedAlgorithms, +} + +impl SkipServerVerification { + pub fn new() -> Result, AttestedTlsError> { + Ok(Arc::new(Self { + supported_algs: Arc::new( + CryptoProvider::get_default().ok_or(AttestedTlsError::NoCryptoProvider)?, + ) + .clone() + .signature_verification_algorithms, + })) + } +} + +impl rustls::client::danger::ServerCertVerifier for SkipServerVerification { + fn verify_server_cert( + &self, + end_entity: &CertificateDer<'_>, + _intermediates: &[CertificateDer<'_>], + _server_name: &pki_types::ServerName<'_>, + _ocsp_response: &[u8], + _now: pki_types::UnixTime, + ) -> Result { + // Parse the certificate + let (_, cert) = X509Certificate::from_der(end_entity).map_err(|_| { + rustls::Error::InvalidCertificate(rustls::CertificateError::BadEncoding) + })?; + + // Verify signature + cert.verify_signature(None).map_err(|_| { + rustls::Error::InvalidCertificate(rustls::CertificateError::BadSignature) + })?; + + Ok(rustls::client::danger::ServerCertVerified::assertion()) + } + + fn verify_tls12_signature( + &self, + message: &[u8], + cert: &CertificateDer<'_>, + dss: &rustls::DigitallySignedStruct, + ) -> Result { + let provider = rustls::crypto::CryptoProvider::get_default() + .ok_or_else(|| rustls::Error::General("No crypto provider installed".into()))?; + + rustls::crypto::verify_tls12_signature( + message, + cert, + dss, + &provider.signature_verification_algorithms, + )?; + + Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) + } + + fn verify_tls13_signature( + &self, + message: &[u8], + cert: &CertificateDer<'_>, + dss: &rustls::DigitallySignedStruct, + ) -> Result { + let provider = rustls::crypto::CryptoProvider::get_default() + .ok_or_else(|| rustls::Error::General("No crypto provider installed".into()))?; + + rustls::crypto::verify_tls13_signature( + message, + cert, + dss, + &provider.signature_verification_algorithms, + )?; + + Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) + } + + fn supported_verify_schemes(&self) -> Vec { + self.supported_algs.supported_schemes() + } +} + +/// Used to allow verification of self-signed certificates during client authentication +#[derive(Debug)] +pub struct SkipClientVerification { + supported_algs: rustls::crypto::WebPkiSupportedAlgorithms, +} + +impl SkipClientVerification { + pub fn new() -> std::sync::Arc { + std::sync::Arc::new(Self { + supported_algs: Arc::new(CryptoProvider::get_default().unwrap()) + .clone() + .signature_verification_algorithms, + }) + } +} + +impl rustls::server::danger::ClientCertVerifier for SkipClientVerification { + fn verify_client_cert( + &self, + end_entity: &CertificateDer<'_>, + _intermediates: &[CertificateDer], + _now: rustls::pki_types::UnixTime, + ) -> Result { + // Parse the certificate + let (_, cert) = X509Certificate::from_der(end_entity).map_err(|_| { + rustls::Error::InvalidCertificate(rustls::CertificateError::BadEncoding) + })?; + + // Verify signature + cert.verify_signature(None).map_err(|_| { + rustls::Error::InvalidCertificate(rustls::CertificateError::BadSignature) + })?; + Ok(rustls::server::danger::ClientCertVerified::assertion()) + } + + fn root_hint_subjects(&self) -> &[rustls::DistinguishedName] { + &[] + } + + fn verify_tls12_signature( + &self, + message: &[u8], + cert: &CertificateDer<'_>, + dss: &rustls::DigitallySignedStruct, + ) -> Result { + let provider = rustls::crypto::CryptoProvider::get_default() + .ok_or_else(|| rustls::Error::General("No crypto provider installed".into()))?; + + rustls::crypto::verify_tls12_signature( + message, + cert, + dss, + &provider.signature_verification_algorithms, + )?; + + Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) + } + + fn verify_tls13_signature( + &self, + message: &[u8], + cert: &CertificateDer<'_>, + dss: &rustls::DigitallySignedStruct, + ) -> Result { + let provider = rustls::crypto::CryptoProvider::get_default() + .ok_or_else(|| rustls::Error::General("No crypto provider installed".into()))?; + + rustls::crypto::verify_tls13_signature( + message, + cert, + dss, + &provider.signature_verification_algorithms, + )?; + + Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) + } + + fn supported_verify_schemes(&self) -> Vec { + self.supported_algs.supported_schemes() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{ + AttestedTlsClient, AttestedTlsServer, + attestation::{AttestationGenerator, AttestationType, AttestationVerifier}, + test_helpers::{generate_certificate_chain, generate_tls_config}, + }; + use tokio::net::TcpListener; + use tokio_rustls::rustls::pki_types::ServerName; + + #[tokio::test] + async fn self_signed_server_attestation() { + let _ = rustls::crypto::aws_lc_rs::default_provider().install_default(); + + let cert_and_key = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + + let server_config = rustls::ServerConfig::builder() + .with_no_client_auth() + .with_single_cert( + cert_and_key.cert_chain.clone().to_vec(), + cert_and_key.key.clone_key(), + ) + .unwrap(); + + let server = AttestedTlsServer::new_with_tls_config( + cert_and_key.cert_chain, + server_config.into(), + AttestationGenerator::new(AttestationType::DcapTdx, None).unwrap(), + AttestationVerifier::expect_none(), + ) + .unwrap(); + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let server_addr = listener.local_addr().unwrap(); + + tokio::spawn(async move { + let (tcp_stream, _) = listener.accept().await.unwrap(); + let (_stream, _measurements, _attestation_type) = + server.handle_connection(tcp_stream).await.unwrap(); + }); + + let client_config = client_tls_config_allow_self_signed().unwrap(); + + let client = AttestedTlsClient::new_with_tls_config( + client_config.into(), + AttestationGenerator::with_no_attestation(), + AttestationVerifier::mock(), + None, + ) + .unwrap(); + + let (_stream, _measurements, _attestation_type) = + client.connect_tcp(&server_addr.to_string()).await.unwrap(); + } + + #[tokio::test] + async fn nested_tls_with_self_signed_server_attestation() { + // Outer TLS setup + let (cert_chain, private_key) = generate_certificate_chain("127.0.0.1".parse().unwrap()); + let (outer_server_config, outer_client_config) = + generate_tls_config(cert_chain.clone(), private_key); + + let outer_acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(outer_server_config)); + let outer_connector = tokio_rustls::TlsConnector::from(Arc::new(outer_client_config)); + + // Inner TLS setup + let cert_and_key = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + + let server_config = rustls::ServerConfig::builder() + .with_no_client_auth() + .with_single_cert( + cert_and_key.cert_chain.clone().to_vec(), + cert_and_key.key.clone_key(), + ) + .unwrap(); + + let server = AttestedTlsServer::new_with_tls_config( + cert_and_key.cert_chain, + server_config.into(), + AttestationGenerator::new(AttestationType::DcapTdx, None).unwrap(), + AttestationVerifier::expect_none(), + ) + .unwrap(); + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let server_addr = listener.local_addr().unwrap(); + + tokio::spawn(async move { + let (tcp_stream, _) = listener.accept().await.unwrap(); + + // Do outer TLS handshake + let tls_stream = outer_acceptor.accept(tcp_stream).await.unwrap(); + + // Do inner (attested) TLS + let (_stream, _measurements, _attestation_type) = + server.handle_connection(tls_stream).await.unwrap(); + }); + + // Inner TLS config + let client_config = client_tls_config_allow_self_signed().unwrap(); + + let client = AttestedTlsClient::new_with_tls_config( + client_config.into(), + AttestationGenerator::with_no_attestation(), + AttestationVerifier::mock(), + None, + ) + .unwrap(); + + let client_tcp_stream = tokio::net::TcpStream::connect(&server_addr).await.unwrap(); + + // Outer TLS handshake + let server_name = ServerName::try_from(server_addr.ip().to_string()).unwrap(); + let tls_stream = outer_connector + .connect(server_name, client_tcp_stream) + .await + .unwrap(); + + // Inner (attested) TLS + let (_stream, _measurements, _attestation_type) = client + .connect(&server_addr.to_string(), tls_stream) + .await + .unwrap(); + } +} diff --git a/src/lib.rs b/src/lib.rs index 9e2db67..b36c0d9 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -2,7 +2,6 @@ pub mod attested_get; pub mod file_server; pub mod health_check; -pub mod normalize_pem; pub mod self_signed; pub use attested_tls; diff --git a/src/main.rs b/src/main.rs index f4d5283..ba6efaa 100644 --- a/src/main.rs +++ b/src/main.rs @@ -23,7 +23,6 @@ use attested_tls_proxy::{ }, file_server::attested_file_server, get_tls_cert, health_check, - normalize_pem::normalize_private_key_pem_to_pkcs8, }; const GIT_REV: &str = match option_env!("GIT_REV") { @@ -529,8 +528,8 @@ fn load_certs_pem(path: PathBuf) -> std::io::Result> /// load TLS private key from a PEM-encoded file fn load_private_key_pem(path: PathBuf) -> anyhow::Result> { - let pem_bytes = std::fs::read(path)?; - normalize_private_key_pem_to_pkcs8(&pem_bytes) + rustls_pemfile::private_key(&mut std::io::BufReader::new(File::open(path)?))? + .ok_or_else(|| anyhow!("No private key found in PEM")) } /// Given a certificate chain, convert it to a PEM encoded string @@ -562,6 +561,126 @@ fn parse_max_in_flight_requests(value: &str) -> Result { #[cfg(test)] mod tests { use super::*; + use tokio_rustls::rustls::{SignatureScheme, crypto::aws_lc_rs}; + + fn load_pem_fixture(pem: &[u8]) -> anyhow::Result> { + let file = tempfile::NamedTempFile::new()?; + std::fs::write(file.path(), pem)?; + load_private_key_pem(file.path().to_owned()) + } + + // Public test fixtures only; never use these keys in a deployment. + const RSA_PKCS1_PEM: &str = r#" +-----BEGIN RSA PRIVATE KEY----- +MIIEpAIBAAKCAQEAsvbL9Jh5+CRiwD4rdixOmHcI/vpwUD0j8PDpDStTTICGpSqN +l7WlSMaFJn5Tc9aXgSftKDiBnPQzPEBBBVxnbJ8MgIY4YilBehoGBY035CPZ+C8P +8wZN7+VoRATUYPYzFdq/cdyCPqB+ZpnJjIRy5WXDlPO8fuGlx5+IvUwEeVIQAXsE +AkXg0Ky3PnB5gyDinCGTM3eM77SzFuWU5LptZjPa9Aap9/QoCrXkC+sbX7pOsWYe +U8JJErIfqBdBT4s/1tVqRmll2Fljr0O65O348zjqkZiQJvRpWvRtSQHK4VurUhHj +7sO6qmdbXD4P9I9Vrug+pAO1J0YDemKpcnO/WQIDAQABAoIBABDOwuru4w2eBTQ+ +4oAPuzXwgATKan/urhBz379f4UvfCkY6z99+rM4/7sNlu9q2PbZglJJhdDLUcHdp +JXImcoQuD9OGR4dYjpC0Hvqof6ZKg68eZGYTooA0UG2K8pNErBmSWMaNyiGtmxFx +wg8TZWMMAqlblsln0dUEs6frmsP1+3AQ8BKyJFCV2TOipf/ja9TcNu9n6ukSwJml +mmDxJS3gTLWxfB0dQs1V+zgLDvqQqjLlgXRXQ8tIualvYY6+tHNJuxeVhyevarGy +lQ1p7GqNFedKQpqMwaXrI/rMY8q75/C0ajKBO7TJZMPRnD5airTNiZ1VG9J+OQrh +Kshdyc0CgYEA92yozUCyY0ns8qs97ixZc3SuKMA6wcmxZEiLv64M64MdKoBM3wfm +wDQGQodg1T3H3Rzw1fZRbaJ4KweG7TKQuey5CY1j7uZyNVNMbGF12Z5HjJhvC88/ +lIpB44aYgmOerqrQczX8KVak8kttw+DoQYGbEyubJ/LXfGu1NCnFZ7MCgYEAuSq0 +LbRMneV9RVMG4z4Y7MrdXBM1C1NcyUdK5nOhUNlWDSlPKIltxznvHoBT6XsYDYb+ +mwPc6Hm6ui75RBhPMlsmoqIeriumnT1Cbr9nk2VZ0+nEKUN6QuG3qH+j8flnh0vc +39wIJs8I2DuYr5EaiUlIaTLWDrKphk3uOzLyNsMCgYEAhuhyaef63HR0hCSm0fTQ +mUlnpMSbxQpKdRmxSUSHuup0vrXSNFHEmcxEFYZnYB4dmgyrrJ5v682IpD2objEC +BL50bibv9FUmtLjElNvXPF83OAvtkIziaAWyw3KiOYZEAY0Vt5wZ8BhUO+Cw6vr4 +6K7YdW1zXiblI+w+k0CraE0CgYAZEF25PgmM6e5t/tIU2mf3TXJvLy5j7RHHMP5D +eW1hizmpqGjNnOSeLgpe/5HcLcxQsHAwPXKeiTOsVgVpoTy/HTV6mCU9AC2aZRtj +8Eat3e8tzxu9ViPrf7Ajf7uKWm8YEj3Ak4EK98VDt7VwNlz4LlI94yK0dJyb0Fqp +6rh8jwKBgQCRq5Ot6bClmTPzQo1T52BGtHZBhVGqFc/76J0knZHL/Qper6TG5IP9 +L4yOiMqxV7mUGWsDBEAi/sisVSLCZvsdWUFvJ4Bp9YBOSWyZlRK59gsOGzS6Qjw7 +fCW7RLfTr0dg2eU3oUTI4B8SHULOhGjSjQ4KCGnbbdIUW2NlFdDkVQ== +-----END RSA PRIVATE KEY----- +"#; + + const RSA_PKCS8_PEM: &str = r#" +-----BEGIN PRIVATE KEY----- +MIIEvAIBADANBgkqhkiG9w0BAQEFAASCBKYwggSiAgEAAoIBAQChw1TJMP2aYJnY +0wG4ElBAXVyFkQvgsx7Sh7yQjbs/WlSE7VOBOJIwKGhVWJk9vJpRWYl6WhiQn/ui +msd06YPkZhvoIotiESyQI7RRpv4YHWj5n8Gomwphj28ttLx3u7AiUWq9uK7mIsFT +Cf6YzUIYQf6FlQu+mYDObtSOtZWmSW69NtZWi2YyXNHPcooPnol8Y0OOd+V6XZrK +Qq8hGfj/7B6HjGbNUH02sKSWC7H7pn+BglNNpX0Znrx+oeEVH4pycYarolixVC0p +N7aK/v+jWiSc3U3Nz6lWALVuzsl2hOC10ie0kVNgheh4bP78YoglKiLHMVO/CRo6 +IzjXeuWTAgMBAAECggEAGarj5jS22OshHk2FBU8qmrv1tV/pkZL6fg95tTo4Dvpn +VNxPlr6CO8/9liVD0478sZHSha6MHU61X/zNT1jKS9CD9xacJUhyWMDBmP81bGAm +Sw21beqEAC0BSDBYg2strJRcqpQGdI/pOyLn2hkftreqCko3Hdw/mwHtCmP3xfW6 +QrmmSOQVq7hKSVCRSs6Do+SW+BLJLb/7ZoU4V8g8nakGh9oXmVKl5CDn/w9f+NSc +VUatPt2+7GMCnUKmQ9qodcuz/EINkimmZY1L2e9WhF9ETm/l292j9Vo2bU7KNoBB +E+9Cn+wMh23mmacSmY7S9SDBkQgVKmRMAyoFxTPmAQKBgQDXwXZLyhYDrtBMOKED +IgFeXMSQj5JZZ10suXYd3nX7iapiNimm5Febe/b4UinhwTzne6WanSvs92vPR5xN +XbJOcep4+YLFt30ZyAb8tyekkdC13rbKNraFWV1+wBs6yT3lfal5JF7Ko7TezYuG +E+nJdrzndLeV8o0yZy8zxT9FAQKBgQC/76rBFpqxBlHCVLLjPEwKvbvsgtRIrqfg +TeKUPS+QHMezYnrQkOiODyUA0/Xs4NBrDI7XuA6tjG0ZLPJWbjckNcvdka0bFQmN +jXXqnTwBsEcFlUFdXVt3EmuIn0K++EBemLndFhn0Wscwn7AO6cQWTA+qy+xuP0Pj +5gGo30hGkwKBgDk3kQugWB456fuMuQZ/qiVALNC5gnI7OzZ1KKHbMSa353uMKZec +zq7pPSG1iG3aNTCeVdie/dsl8m1R7F2ID5VGGIxkfw24D3Ea3t9+IwE9uj/BBHCz ++ct7W5QVliMM42FM5fi+cHUE3R6JHAs+lK1c09P93AHkBRXsz1PHZ3QBAoGAbWgD +MG9fHBtbDWfUVH0xZ0oBze5BbXDJVq1uw0shSodtOg6frTV8qkVttUwdObpocyzE +W6iaDUkngxtAxA2tNuHHZHQ+dVqHiH2jQmoAI4JE6aTLjpnBol0ImOcXV94Qaxup +jqGjh8sbEddktwt/b6pJn/T/v1QmsciRF563BysCgYB7t1HngCHE6zCGQdxDDp5A +Vb9rJaarKKy5TIX0svK4iGxmoD/qCf9o5LwKwoQLPf4K2F8fhEZZGes+62M5HzmZ +FCKYluqG2/M/gs/AE9K+btrpuIbZZB5Prris+THkBBxHTt49WFxwkVK+CbQxg0VC +32K3vJkhe2O33oHoyzQRfw== +-----END PRIVATE KEY----- +"#; + + const P256_SEC1_PEM: &str = r#"-----BEGIN EC PRIVATE KEY----- +MHcCAQEEIP37GKC//8GKtvmYmf62bpDsD8vlhlxLZ1PNbTICsvo9oAoGCCqGSM49 +AwEHoUQDQgAEoCwAV5jHuPli5xYkmgQiGsa+MsZLXXmqrUR5Wu0S5Xgsm5lv/wy3 +JSUC8mADyuZZsVyaFkSgSGkyyJfwSvVLNg== +-----END EC PRIVATE KEY----- +"#; + + #[test] + fn original_key_formats_work_with_tls_provider() { + let provider = aws_lc_rs::default_provider(); + for (pem, scheme) in [ + (RSA_PKCS1_PEM, SignatureScheme::RSA_PSS_SHA256), + (RSA_PKCS8_PEM, SignatureScheme::RSA_PSS_SHA256), + (P256_SEC1_PEM, SignatureScheme::ECDSA_NISTP256_SHA256), + ] { + let key = load_pem_fixture(pem.as_bytes()).unwrap(); + match pem { + RSA_PKCS1_PEM => assert!(matches!(&key, PrivateKeyDer::Pkcs1(_))), + RSA_PKCS8_PEM => assert!(matches!(&key, PrivateKeyDer::Pkcs8(_))), + P256_SEC1_PEM => assert!(matches!(&key, PrivateKeyDer::Sec1(_))), + _ => unreachable!(), + } + let signing_key = provider.key_provider.load_private_key(key).unwrap(); + let signer = signing_key.choose_scheme(&[scheme]).unwrap(); + assert!( + !signer + .sign(b"TLS key loading regression test") + .unwrap() + .is_empty() + ); + } + } + + #[test] + fn missing_and_malformed_keys_fail_and_other_pem_blocks_are_skipped() { + let certificate = "-----BEGIN CERTIFICATE-----\nAA==\n-----END CERTIFICATE-----\n"; + for pem in [ + "", + "not PEM", + certificate, + "-----BEGIN PRIVATE KEY-----\ninvalid base64!\n-----END PRIVATE KEY-----\n", + ] { + assert!(load_pem_fixture(pem.as_bytes()).is_err()); + } + let bundle = format!("{certificate}{RSA_PKCS1_PEM}{P256_SEC1_PEM}"); + assert!(matches!( + load_pem_fixture(bundle.as_bytes()).unwrap(), + PrivateKeyDer::Pkcs1(_) + )); + } /// Checks that CLI parsing rejects invalid limits before starting the proxy. #[test] diff --git a/src/normalize_pem.rs b/src/normalize_pem.rs deleted file mode 100644 index 810e9d4..0000000 --- a/src/normalize_pem.rs +++ /dev/null @@ -1,134 +0,0 @@ -use anyhow::{Result, anyhow, bail}; -use pkcs8::EncodePrivateKey; -use std::io::Cursor; -use tokio_rustls::rustls::pki_types::{PrivateKeyDer, PrivatePkcs8KeyDer}; - -/// Given a PEM encoded private key convert to PKCS8 which Rustls accepts -pub fn normalize_private_key_pem_to_pkcs8(pem: &[u8]) -> Result> { - let der = normalize_private_key_pem_to_pkcs8_der(pem)?; - let pkcs8_der = PrivatePkcs8KeyDer::from(der); - Ok(pkcs8_der.into()) -} - -fn normalize_private_key_pem_to_pkcs8_der(pem: &[u8]) -> Result> { - let mut rd = Cursor::new(pem); - - // Find first private key in the PEM (ignore certs, etc.) - let item = loop { - match rustls_pemfile::read_one(&mut rd).map_err(|e| anyhow!("reading PEM: {e}"))? { - Some(it) => match it { - rustls_pemfile::Item::Pkcs8Key(_) - | rustls_pemfile::Item::Pkcs1Key(_) - | rustls_pemfile::Item::Sec1Key(_) => break it, - _ => continue, - }, - None => bail!("No private key found in PEM"), - } - }; - - match item { - // Already PKCS#8: pass through DER bytes - rustls_pemfile::Item::Pkcs8Key(k) => Ok(k.secret_pkcs8_der().to_vec()), - - // RSA PKCS#1 ("BEGIN RSA PRIVATE KEY") -> PKCS#8 - rustls_pemfile::Item::Pkcs1Key(k) => { - use pkcs1::DecodeRsaPrivateKey; - use rsa::RsaPrivateKey; - - let key = RsaPrivateKey::from_pkcs1_der(k.secret_pkcs1_der()) - .map_err(|e| anyhow!("Parsing PKCS#1 RSA key: {e:?}"))?; - - let pkcs8 = key - .to_pkcs8_der() - .map_err(|e| anyhow!("Encoding PKCS#8 RSA key: {e:?}"))?; - - Ok(pkcs8.as_bytes().to_vec()) - } - - // SEC1 ("BEGIN EC PRIVATE KEY") for P-256 -> PKCS#8 - rustls_pemfile::Item::Sec1Key(k) => { - let sk = p256::SecretKey::from_sec1_der(k.secret_sec1_der()) - .map_err(|e| anyhow!("Parsing SEC1 P-256 key: {e:?}"))?; - - let pkcs8 = sk - .to_pkcs8_der() - .map_err(|e| anyhow!("Encoding PKCS#8 P-256 key: {e:?}"))?; - - Ok(pkcs8.as_bytes().to_vec()) - } - - _ => Err(anyhow!("unexpected PEM item (filtered earlier)")), - } -} - -#[cfg(test)] -mod tests { - use super::*; - - const RSA_PKCS1_PEM: &str = r#" ------BEGIN RSA PRIVATE KEY----- -MIIEpAIBAAKCAQEAsvbL9Jh5+CRiwD4rdixOmHcI/vpwUD0j8PDpDStTTICGpSqN -l7WlSMaFJn5Tc9aXgSftKDiBnPQzPEBBBVxnbJ8MgIY4YilBehoGBY035CPZ+C8P -8wZN7+VoRATUYPYzFdq/cdyCPqB+ZpnJjIRy5WXDlPO8fuGlx5+IvUwEeVIQAXsE -AkXg0Ky3PnB5gyDinCGTM3eM77SzFuWU5LptZjPa9Aap9/QoCrXkC+sbX7pOsWYe -U8JJErIfqBdBT4s/1tVqRmll2Fljr0O65O348zjqkZiQJvRpWvRtSQHK4VurUhHj -7sO6qmdbXD4P9I9Vrug+pAO1J0YDemKpcnO/WQIDAQABAoIBABDOwuru4w2eBTQ+ -4oAPuzXwgATKan/urhBz379f4UvfCkY6z99+rM4/7sNlu9q2PbZglJJhdDLUcHdp -JXImcoQuD9OGR4dYjpC0Hvqof6ZKg68eZGYTooA0UG2K8pNErBmSWMaNyiGtmxFx -wg8TZWMMAqlblsln0dUEs6frmsP1+3AQ8BKyJFCV2TOipf/ja9TcNu9n6ukSwJml -mmDxJS3gTLWxfB0dQs1V+zgLDvqQqjLlgXRXQ8tIualvYY6+tHNJuxeVhyevarGy -lQ1p7GqNFedKQpqMwaXrI/rMY8q75/C0ajKBO7TJZMPRnD5airTNiZ1VG9J+OQrh -Kshdyc0CgYEA92yozUCyY0ns8qs97ixZc3SuKMA6wcmxZEiLv64M64MdKoBM3wfm -wDQGQodg1T3H3Rzw1fZRbaJ4KweG7TKQuey5CY1j7uZyNVNMbGF12Z5HjJhvC88/ -lIpB44aYgmOerqrQczX8KVak8kttw+DoQYGbEyubJ/LXfGu1NCnFZ7MCgYEAuSq0 -LbRMneV9RVMG4z4Y7MrdXBM1C1NcyUdK5nOhUNlWDSlPKIltxznvHoBT6XsYDYb+ -mwPc6Hm6ui75RBhPMlsmoqIeriumnT1Cbr9nk2VZ0+nEKUN6QuG3qH+j8flnh0vc -39wIJs8I2DuYr5EaiUlIaTLWDrKphk3uOzLyNsMCgYEAhuhyaef63HR0hCSm0fTQ -mUlnpMSbxQpKdRmxSUSHuup0vrXSNFHEmcxEFYZnYB4dmgyrrJ5v682IpD2objEC -BL50bibv9FUmtLjElNvXPF83OAvtkIziaAWyw3KiOYZEAY0Vt5wZ8BhUO+Cw6vr4 -6K7YdW1zXiblI+w+k0CraE0CgYAZEF25PgmM6e5t/tIU2mf3TXJvLy5j7RHHMP5D -eW1hizmpqGjNnOSeLgpe/5HcLcxQsHAwPXKeiTOsVgVpoTy/HTV6mCU9AC2aZRtj -8Eat3e8tzxu9ViPrf7Ajf7uKWm8YEj3Ak4EK98VDt7VwNlz4LlI94yK0dJyb0Fqp -6rh8jwKBgQCRq5Ot6bClmTPzQo1T52BGtHZBhVGqFc/76J0knZHL/Qper6TG5IP9 -L4yOiMqxV7mUGWsDBEAi/sisVSLCZvsdWUFvJ4Bp9YBOSWyZlRK59gsOGzS6Qjw7 -fCW7RLfTr0dg2eU3oUTI4B8SHULOhGjSjQ4KCGnbbdIUW2NlFdDkVQ== ------END RSA PRIVATE KEY----- -"#; - - const RSA_PKCS8_PEM: &str = r#" ------BEGIN PRIVATE KEY----- -MIIEvAIBADANBgkqhkiG9w0BAQEFAASCBKYwggSiAgEAAoIBAQChw1TJMP2aYJnY -0wG4ElBAXVyFkQvgsx7Sh7yQjbs/WlSE7VOBOJIwKGhVWJk9vJpRWYl6WhiQn/ui -msd06YPkZhvoIotiESyQI7RRpv4YHWj5n8Gomwphj28ttLx3u7AiUWq9uK7mIsFT -Cf6YzUIYQf6FlQu+mYDObtSOtZWmSW69NtZWi2YyXNHPcooPnol8Y0OOd+V6XZrK -Qq8hGfj/7B6HjGbNUH02sKSWC7H7pn+BglNNpX0Znrx+oeEVH4pycYarolixVC0p -N7aK/v+jWiSc3U3Nz6lWALVuzsl2hOC10ie0kVNgheh4bP78YoglKiLHMVO/CRo6 -IzjXeuWTAgMBAAECggEAGarj5jS22OshHk2FBU8qmrv1tV/pkZL6fg95tTo4Dvpn -VNxPlr6CO8/9liVD0478sZHSha6MHU61X/zNT1jKS9CD9xacJUhyWMDBmP81bGAm -Sw21beqEAC0BSDBYg2strJRcqpQGdI/pOyLn2hkftreqCko3Hdw/mwHtCmP3xfW6 -QrmmSOQVq7hKSVCRSs6Do+SW+BLJLb/7ZoU4V8g8nakGh9oXmVKl5CDn/w9f+NSc -VUatPt2+7GMCnUKmQ9qodcuz/EINkimmZY1L2e9WhF9ETm/l292j9Vo2bU7KNoBB -E+9Cn+wMh23mmacSmY7S9SDBkQgVKmRMAyoFxTPmAQKBgQDXwXZLyhYDrtBMOKED -IgFeXMSQj5JZZ10suXYd3nX7iapiNimm5Febe/b4UinhwTzne6WanSvs92vPR5xN -XbJOcep4+YLFt30ZyAb8tyekkdC13rbKNraFWV1+wBs6yT3lfal5JF7Ko7TezYuG -E+nJdrzndLeV8o0yZy8zxT9FAQKBgQC/76rBFpqxBlHCVLLjPEwKvbvsgtRIrqfg -TeKUPS+QHMezYnrQkOiODyUA0/Xs4NBrDI7XuA6tjG0ZLPJWbjckNcvdka0bFQmN -jXXqnTwBsEcFlUFdXVt3EmuIn0K++EBemLndFhn0Wscwn7AO6cQWTA+qy+xuP0Pj -5gGo30hGkwKBgDk3kQugWB456fuMuQZ/qiVALNC5gnI7OzZ1KKHbMSa353uMKZec -zq7pPSG1iG3aNTCeVdie/dsl8m1R7F2ID5VGGIxkfw24D3Ea3t9+IwE9uj/BBHCz -+ct7W5QVliMM42FM5fi+cHUE3R6JHAs+lK1c09P93AHkBRXsz1PHZ3QBAoGAbWgD -MG9fHBtbDWfUVH0xZ0oBze5BbXDJVq1uw0shSodtOg6frTV8qkVttUwdObpocyzE -W6iaDUkngxtAxA2tNuHHZHQ+dVqHiH2jQmoAI4JE6aTLjpnBol0ImOcXV94Qaxup -jqGjh8sbEddktwt/b6pJn/T/v1QmsciRF563BysCgYB7t1HngCHE6zCGQdxDDp5A -Vb9rJaarKKy5TIX0svK4iGxmoD/qCf9o5LwKwoQLPf4K2F8fhEZZGes+62M5HzmZ -FCKYluqG2/M/gs/AE9K+btrpuIbZZB5Prris+THkBBxHTt49WFxwkVK+CbQxg0VC -32K3vJkhe2O33oHoyzQRfw== ------END PRIVATE KEY----- -"#; - - #[test] - fn convert_private_key_to_pkcs8() { - let _key = normalize_private_key_pem_to_pkcs8(RSA_PKCS1_PEM.as_bytes()).unwrap(); - let _key = normalize_private_key_pem_to_pkcs8(RSA_PKCS8_PEM.as_bytes()).unwrap(); - } -} diff --git a/src/self_signed.rs b/src/self_signed.rs index 309803f..dc6af95 100644 --- a/src/self_signed.rs +++ b/src/self_signed.rs @@ -1,325 +1,2 @@ -use std::{net::IpAddr, sync::Arc}; -use tokio_rustls::rustls::{ - self, - crypto::CryptoProvider, - pki_types::{self, CertificateDer, PrivatePkcs8KeyDer}, -}; -use x509_parser::prelude::{FromDer, X509Certificate}; - -use crate::attested_tls::{AttestedTlsError, TlsCertAndKey}; - -/// Generate a self signed certifcate -pub fn generate_self_signed_cert(ip_address: IpAddr) -> Result { - let keypair = rcgen::KeyPair::generate()?; - let mut params = rcgen::CertificateParams::default(); - params - .subject_alt_names - .push(rcgen::SanType::IpAddress(ip_address)); - - let cert = params.self_signed(&keypair)?; - Ok(TlsCertAndKey { - cert_chain: vec![cert.der().clone()], - key: PrivatePkcs8KeyDer::from(keypair.serialize_der()).into(), - }) -} - -/// Client TLS configuration which accepts self-signed remote certificates -pub fn client_tls_config_allow_self_signed() -> Result { - Ok(rustls::ClientConfig::builder() - .dangerous() - .with_custom_certificate_verifier(SkipServerVerification::new()?) - .with_no_client_auth()) -} - -/// Used to allow verification of self-signed certificates -#[derive(Debug, Clone)] -pub struct SkipServerVerification { - supported_algs: rustls::crypto::WebPkiSupportedAlgorithms, -} - -impl SkipServerVerification { - pub fn new() -> Result, AttestedTlsError> { - Ok(Arc::new(Self { - supported_algs: Arc::new( - CryptoProvider::get_default().ok_or(AttestedTlsError::NoCryptoProvider)?, - ) - .clone() - .signature_verification_algorithms, - })) - } -} - -impl rustls::client::danger::ServerCertVerifier for SkipServerVerification { - fn verify_server_cert( - &self, - end_entity: &CertificateDer<'_>, - _intermediates: &[CertificateDer<'_>], - _server_name: &pki_types::ServerName<'_>, - _ocsp_response: &[u8], - _now: pki_types::UnixTime, - ) -> Result { - // Parse the certificate - let (_, cert) = X509Certificate::from_der(end_entity).map_err(|_| { - rustls::Error::InvalidCertificate(rustls::CertificateError::BadEncoding) - })?; - - // Verify signature - cert.verify_signature(None).map_err(|_| { - rustls::Error::InvalidCertificate(rustls::CertificateError::BadSignature) - })?; - - Ok(rustls::client::danger::ServerCertVerified::assertion()) - } - - fn verify_tls12_signature( - &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, - ) -> Result { - let provider = rustls::crypto::CryptoProvider::get_default() - .ok_or_else(|| rustls::Error::General("No crypto provider installed".into()))?; - - rustls::crypto::verify_tls12_signature( - message, - cert, - dss, - &provider.signature_verification_algorithms, - )?; - - Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) - } - - fn verify_tls13_signature( - &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, - ) -> Result { - let provider = rustls::crypto::CryptoProvider::get_default() - .ok_or_else(|| rustls::Error::General("No crypto provider installed".into()))?; - - rustls::crypto::verify_tls13_signature( - message, - cert, - dss, - &provider.signature_verification_algorithms, - )?; - - Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) - } - - fn supported_verify_schemes(&self) -> Vec { - self.supported_algs.supported_schemes() - } -} - -/// Used to allow verification of self-signed certificates during client authentication -#[derive(Debug)] -pub struct SkipClientVerification { - supported_algs: rustls::crypto::WebPkiSupportedAlgorithms, -} - -impl SkipClientVerification { - pub fn new() -> std::sync::Arc { - std::sync::Arc::new(Self { - supported_algs: Arc::new(CryptoProvider::get_default().unwrap()) - .clone() - .signature_verification_algorithms, - }) - } -} - -impl rustls::server::danger::ClientCertVerifier for SkipClientVerification { - fn verify_client_cert( - &self, - end_entity: &CertificateDer<'_>, - _intermediates: &[CertificateDer], - _now: rustls::pki_types::UnixTime, - ) -> Result { - // Parse the certificate - let (_, cert) = X509Certificate::from_der(end_entity).map_err(|_| { - rustls::Error::InvalidCertificate(rustls::CertificateError::BadEncoding) - })?; - - // Verify signature - cert.verify_signature(None).map_err(|_| { - rustls::Error::InvalidCertificate(rustls::CertificateError::BadSignature) - })?; - Ok(rustls::server::danger::ClientCertVerified::assertion()) - } - - fn root_hint_subjects(&self) -> &[rustls::DistinguishedName] { - &[] - } - - fn verify_tls12_signature( - &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, - ) -> Result { - let provider = rustls::crypto::CryptoProvider::get_default() - .ok_or_else(|| rustls::Error::General("No crypto provider installed".into()))?; - - rustls::crypto::verify_tls12_signature( - message, - cert, - dss, - &provider.signature_verification_algorithms, - )?; - - Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) - } - - fn verify_tls13_signature( - &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, - ) -> Result { - let provider = rustls::crypto::CryptoProvider::get_default() - .ok_or_else(|| rustls::Error::General("No crypto provider installed".into()))?; - - rustls::crypto::verify_tls13_signature( - message, - cert, - dss, - &provider.signature_verification_algorithms, - )?; - - Ok(rustls::client::danger::HandshakeSignatureValid::assertion()) - } - - fn supported_verify_schemes(&self) -> Vec { - self.supported_algs.supported_schemes() - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::{ - AttestationGenerator, - attestation::{AttestationType, AttestationVerifier}, - attested_tls::{AttestedTlsClient, AttestedTlsServer}, - test_helpers::{generate_certificate_chain, generate_tls_config}, - }; - use tokio::net::TcpListener; - use tokio_rustls::rustls::pki_types::ServerName; - - #[tokio::test] - async fn self_signed_server_attestation() { - let _ = rustls::crypto::aws_lc_rs::default_provider().install_default(); - - let cert_and_key = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); - - let server_config = rustls::ServerConfig::builder() - .with_no_client_auth() - .with_single_cert( - cert_and_key.cert_chain.clone().to_vec(), - cert_and_key.key.clone_key(), - ) - .unwrap(); - - let server = AttestedTlsServer::new_with_tls_config( - cert_and_key.cert_chain, - server_config.into(), - AttestationGenerator::new(AttestationType::DcapTdx, None).unwrap(), - AttestationVerifier::expect_none(), - ) - .unwrap(); - - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let server_addr = listener.local_addr().unwrap(); - - tokio::spawn(async move { - let (tcp_stream, _) = listener.accept().await.unwrap(); - let (_stream, _measurements, _attestation_type) = - server.handle_connection(tcp_stream).await.unwrap(); - }); - - let client_config = client_tls_config_allow_self_signed().unwrap(); - - let client = AttestedTlsClient::new_with_tls_config( - client_config.into(), - AttestationGenerator::with_no_attestation(), - AttestationVerifier::mock(), - None, - ) - .unwrap(); - - let (_stream, _measurements, _attestation_type) = - client.connect_tcp(&server_addr.to_string()).await.unwrap(); - } - - #[tokio::test] - async fn nested_tls_with_self_signed_server_attestation() { - // Outer TLS setup - let (cert_chain, private_key) = generate_certificate_chain("127.0.0.1".parse().unwrap()); - let (outer_server_config, outer_client_config) = - generate_tls_config(cert_chain.clone(), private_key); - - let outer_acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(outer_server_config)); - let outer_connector = tokio_rustls::TlsConnector::from(Arc::new(outer_client_config)); - - // Inner TLS setup - let cert_and_key = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); - - let server_config = rustls::ServerConfig::builder() - .with_no_client_auth() - .with_single_cert( - cert_and_key.cert_chain.clone().to_vec(), - cert_and_key.key.clone_key(), - ) - .unwrap(); - - let server = AttestedTlsServer::new_with_tls_config( - cert_and_key.cert_chain, - server_config.into(), - AttestationGenerator::new(AttestationType::DcapTdx, None).unwrap(), - AttestationVerifier::expect_none(), - ) - .unwrap(); - - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let server_addr = listener.local_addr().unwrap(); - - tokio::spawn(async move { - let (tcp_stream, _) = listener.accept().await.unwrap(); - - // Do outer TLS handshake - let tls_stream = outer_acceptor.accept(tcp_stream).await.unwrap(); - - // Do inner (attested) TLS - let (_stream, _measurements, _attestation_type) = - server.handle_connection(tls_stream).await.unwrap(); - }); - - // Inner TLS config - let client_config = client_tls_config_allow_self_signed().unwrap(); - - let client = AttestedTlsClient::new_with_tls_config( - client_config.into(), - AttestationGenerator::with_no_attestation(), - AttestationVerifier::mock(), - None, - ) - .unwrap(); - - let client_tcp_stream = tokio::net::TcpStream::connect(&server_addr).await.unwrap(); - - // Outer TLS handshake - let server_name = ServerName::try_from(server_addr.ip().to_string()).unwrap(); - let tls_stream = outer_connector - .connect(server_name, client_tcp_stream) - .await - .unwrap(); - - // Inner (attested) TLS - let (_stream, _measurements, _attestation_type) = client - .connect(&server_addr.to_string(), tls_stream) - .await - .unwrap(); - } -} +//! Shared self-signed TLS certificate helpers. +pub use attested_tls::self_signed::*; diff --git a/tcp-tunnel/Cargo.toml b/tcp-tunnel/Cargo.toml new file mode 100644 index 0000000..4e41147 --- /dev/null +++ b/tcp-tunnel/Cargo.toml @@ -0,0 +1,33 @@ +[package] +name = "attested-tls-tcp-tunnel" +version = "0.1.0" +edition = "2024" +license = "MIT" +description = "A remote-attested TLS TCP tunnel" +repository = "https://github.com/flashbots/attested-tls-proxy" +build = "../build.rs" + +[dependencies] +attested-tls = { workspace = true, features = ["self-signed"] } +tokio = { workspace = true, features = ["fs", "io-util", "macros", "net", "rt-multi-thread", "signal", "sync", "time"] } +tokio-rustls = { workspace = true, features = ["aws_lc_rs"] } +thiserror.workspace = true +anyhow.workspace = true +clap.workspace = true +rustls-pemfile.workspace = true +serde_json.workspace = true +webpki-roots.workspace = true +tracing.workspace = true +tracing-subscriber.workspace = true + +[features] +default = [] +azure = ["attested-tls/azure"] + +[dev-dependencies] +attested-tls = { workspace = true, features = ["test-helpers", "mock"] } +tokio = { workspace = true, features = ["process"] } +tempfile.workspace = true +h2.workspace = true +http.workspace = true +bytes.workspace = true diff --git a/tcp-tunnel/README.md b/tcp-tunnel/README.md new file mode 100644 index 0000000..8cdb3ad --- /dev/null +++ b/tcp-tunnel/README.md @@ -0,0 +1,184 @@ +# attested-tls-tcp-tunnel + +An attested-TLS TCP tunnel with a `client` and `server` CLI and library. +It forwards bytes without parsing any application protocol: + +```text +source <-- TCP --> tunnel client <-- attested TLS TCP --> tunnel server <-- TCP --> target +``` + +Each source TCP connection opens a dedicated attested-TLS connection and target +TCP connection. Separate source connections get separate tunnels; no connections are +pooled. The server forwards to one configured target, which may itself dispatch +requests to a pool of workers. + +## Local example + +Build the executable: + +```sh +cargo build -p attested-tls-tcp-tunnel +``` + +For a quick round trip, start a local HTTP service in one terminal: + +```sh +python3 -m http.server 8000 --bind 127.0.0.1 +``` + +Start the tunnel server in another terminal. With no certificate/key arguments, +it generates a self-signed certificate. This example explicitly disables local +attestation and accepts a client with no attestation: + +```sh +target/debug/attested-tls-tcp-tunnel server \ + --listen-addr 127.0.0.1:7000 \ + --server-attestation-type none \ + --allowed-remote-attestation-type none \ + 127.0.0.1:8000 +``` + +Start the tunnel client in another terminal: + +```sh +target/debug/attested-tls-tcp-tunnel client \ + --listen-addr 127.0.0.1:6000 \ + --client-attestation-type none \ + --allowed-remote-attestation-type none \ + --allow-self-signed \ + 127.0.0.1:7000 +``` + +Then `curl http://127.0.0.1:6000/` reaches the target. These `none` policies are for +demonstrating transport without a CVM. For attested deployments, select the local +attestation type (or leave automatic detection enabled) and configure accepted +remote measurements using `--measurements-file`. + +For gRPC, point the tunnel server at your gRPC service instead of port 8000 and +configure the gRPC client to use plaintext HTTP/2 at `127.0.0.1:6000`. Unary and all +streaming RPC types use the same forwarding path. RPC metadata, trailers, +cancellation, PINGs, and GOAWAY travel between the actual gRPC endpoints. + +Application TLS also works: configure TLS in the gRPC client and target as usual. +The inner TLS session passes through unchanged. Configure the application's +server name/authority for its real service identity rather than the local tunnel +address. The tunnel's `--tls-*` settings configure only the outer attested TLS. + +## Configuration + +Both commands require exactly one of `--measurements-file ` and +`--allowed-remote-attestation-type `. Measurement policy and attestation +settings follow the [HTTP proxy CLI](../README.md#measurements-file). + +| Option | Default / behavior | +|---|---| +| `--listen-addr`, `-l` (`LISTEN_ADDR`) | Client `127.0.0.1:0`; server `0.0.0.0:0`. The actual bound address is logged. | +| Positional target | Client `host[:port]`, default port 443; server `host:port` with required port. IPv6 literals use brackets. | +| `--setup-timeout-secs` | 60; includes DNS, connect, TLS, attestation, and the server's target connect. Applied independently at each endpoint. | +| `--max-connections` | 256 per listener, including connections still establishing; excess arrivals are immediately closed. | +| `--shutdown-grace-secs` | 30; stop accepting, drain, then close remaining tunnels. Zero closes immediately. | +| `--tls-private-key-path`, `--tls-certificate-path` | Must be supplied together; accept PKCS#8, RSA PKCS#1, and P-256 SEC1 PEM keys. | +| Client `--tls-ca-certificate` | Trust the first PEM certificate instead of public roots. | +| Client `--allow-self-signed` | Accept a self-signed server certificate; still verify attestation. Cannot combine with `--tls-ca-certificate`. Preserves any supplied client identity. | +| Server `--client-auth` | Require a TLS client certificate authenticated against public roots. Private client CAs can be configured through the Rust API. | +| `--client-attestation-type` / `--server-attestation-type` | Automatic local detection when omitted. | +| `--pccs-url`, `--dev-dummy-dcap` | Same meaning as in the HTTP proxy. | +| `--log-debug`, `--log-json`, `--log-dcap-quote` | Debug logs, structured logs, or DCAP quote dumps in `quotes/`. Payload bytes are never logged. | +| `--override-azure-outdated-tcb` | Same Azure verification override as in the HTTP proxy. | + +The existing environment names are supported: `MEASUREMENTS_FILE`, +`TLS_PRIVATE_KEY_PATH`, `TLS_CERTIFICATE_PATH`, `CLIENT_ATTESTATION_TYPE`, +`SERVER_ATTESTATION_TYPE`, and `OVERRIDE_AZURE_OUTDATED_TCB`. + +Enable Azure support with `--features azure`; this requires the same TPM system +dependencies as the HTTP proxy. There is no health-check listener in this version. + +## Lifecycle and trust + +CLI startup binds the listener without contacting the remote service. Each accepted +connection gets one setup attempt. Failures close that connection and are logged +with the endpoint and phase; the tunnel does not send HTTP/gRPC error messages. +It does not retry, reconnect an established stream, or replay application data. +gRPC channels and applications own reconnection, retries, and stream recovery. + +The client verifies the remote server before forwarding source bytes. The server +verifies the client and application protocol before connecting to its target. +The negotiated ALPN must be `flashbots-ratls/1+tcp-tunnel`; the base transport's +bare ALPN fallback is rejected. An HTTP proxy is not a compatible tunnel peer. +After the existing attestation exchange, there is no additional tunnel framing +or readiness acknowledgement. Client-side setup completion does not confirm +that the remote target has connected; a subsequent failure closes the stream. + +The local TCP legs rely on their deployment's trust boundary unless the +application supplies its own TLS. Attestation is checked once per new tunnel; +it does not continuously re-attest a long-lived connection or identify each +application sharing a local proxy. The target sees the tunnel server's IP and +receives no injected attestation metadata. + +Forwarding uses bounded buffers and preserves half-closes: a sender may finish +its upload while still receiving a response. There is no established-stream +idle timeout or lifetime limit. Configure RPC deadlines and keepalive in gRPC. +A transport error closes the affected tunnel, including its concurrent RPCs. + +SIGINT/Ctrl-C and SIGTERM stop acceptance and trigger bounded draining. An opaque +tunnel cannot generate gRPC GOAWAY or drain individual RPCs. Long-lived streams +are interrupted if still active when the grace period expires. Blocking quote +generation cannot be canceled by dropping its async task; the executable uses +bounded runtime shutdown after draining sockets. The connection cap bounds live +connections, not quote-generation work that outlives a setup timeout. + +## Rust API + +`TunnelClient` and `TunnelServer` offer `new`, `new_with_tls_config`, `local_addr`, +and `serve_until`. `TunnelOptions` contains setup timeout, connection count, and +shutdown grace. The custom constructors accept Rustls configurations alongside +attestation generators/verifiers and matching certificate chains. They replace +ALPN with the tunnel protocol. Initialize a Rustls crypto provider before use. + +```rust,no_run +use attested_tls::attestation::{AttestationGenerator, AttestationVerifier}; +use attested_tls_tcp_tunnel::{TunnelClient, TunnelOptions}; + +# async fn example() -> Result<(), Box> { +let _ = tokio_rustls::rustls::crypto::aws_lc_rs::default_provider().install_default(); +let client = TunnelClient::new( + "127.0.0.1:6000", + "tunnel.example.com:443".into(), + None, // optional client TLS identity + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), // replace with your measurement policy + None, // use public CA roots + false, // No startup check. + TunnelOptions::default(), +).await?; +client.serve_until(async { + let _ = tokio::signal::ctrl_c().await; +}).await?; +# Ok(()) +# } +``` + +Dropping a `serve_until` future aborts its connection tasks. For graceful shutdown, +resolve the supplied shutdown future and await completion instead. Embedding +applications own runtime shutdown, including outstanding blocking attestation +work. The `tls` module provides TLS configuration helpers, including self-signed +verification that retains client credentials. + +Both client constructors accept a `startup_check` boolean before `options`. +When true, construction binds the listener, verifies an upstream connection +(TLS, attestation, and tunnel ALPN), and closes that probe before returning. +The check uses `setup_timeout`; a failure returns an error and drops the listener. +It may cause an empty connection to the server's target, but does not confirm +target connectivity because the protocol has no readiness acknowledgement. +The CLI passes false and has no startup-check option. + +## Validation + +```sh +cargo check -p attested-tls-tcp-tunnel +cargo test -p attested-tls-tcp-tunnel --all-targets +``` + +Tests use mock attestation only as a development dependency. The HTTP/2 fixture +tests gRPC message framing, metadata and status trailers, multiplexed streaming, +and cancellation under response flow control without a protobuf compiler. diff --git a/tcp-tunnel/src/lib.rs b/tcp-tunnel/src/lib.rs new file mode 100644 index 0000000..2ea2eea --- /dev/null +++ b/tcp-tunnel/src/lib.rs @@ -0,0 +1,521 @@ +//! Tunnel a TCP connection over remote attested TLS. +//! +//! This is a proxy which accepts TCP connections and does TLS handshake followed +//! by an attestation exchange. The TLS session byte-stream is then handed back to +//! the calling application. +//! +//! Assumes a Rustls crypto provider is already installed. +pub mod tls; + +use std::{future::Future, io, net::SocketAddr, num::NonZeroUsize, sync::Arc, time::Duration}; + +use attested_tls::{ + AttestedTlsClient, AttestedTlsError, AttestedTlsServer, SUPPORTED_ALPN_PROTOCOL_VERSIONS, + TlsCertAndKey, + attestation::{AttestationGenerator, AttestationVerifier}, +}; +use tokio::{ + io::AsyncWriteExt, + net::{TcpListener, TcpStream, ToSocketAddrs}, + sync::Semaphore, + task::JoinSet, +}; +use tokio_rustls::rustls::{ClientConfig, ServerConfig, pki_types::CertificateDer}; +use tracing::Instrument; + +const APPLICATION_PROTOCOL: &[u8] = b"tcp-tunnel"; + +/// Limits shared by the client and server. No limit applies to the lifetime or +/// inactivity of an established tunnel. +#[derive(Clone, Copy, Debug)] +pub struct TunnelOptions { + /// Timeout for connections establishment, TLS handshake and attestation exchange + pub setup_timeout: Duration, + /// Counts both connections being established and established connections. + pub max_connections: NonZeroUsize, + /// Timeout for closing connections during graceful shutdown. + pub shutdown_grace: Duration, +} + +impl Default for TunnelOptions { + fn default() -> Self { + Self { + setup_timeout: Duration::from_secs(60), + max_connections: NonZeroUsize::new(256).unwrap(), + shutdown_grace: Duration::from_secs(30), + } + } +} + +impl TunnelOptions { + fn validate(self) -> Result { + if self.setup_timeout.is_zero() { + return Err(TunnelError::Configuration("setup timeout must be positive")); + } + if self.max_connections.get() > Semaphore::MAX_PERMITS { + return Err(TunnelError::Configuration( + "connection limit exceeds Tokio's maximum", + )); + } + Ok(self) + } +} + +/// Accepts local TCP connections and opens a dedicated attested connection for +/// each. Constructors contact the server only when `startup_check` is enabled. +pub struct TunnelClient(Tunnel); + +impl TunnelClient { + /// If `startup_check` is true, verify an upstream connection within + /// `options.setup_timeout` and close it before returning. This checks TLS, + /// attestation, and ALPN, not final target reachability. The server may open + /// an empty target connection. On failure the local listener is dropped. + #[allow(clippy::too_many_arguments)] // Keep the startup check explicit in the constructor. + pub async fn new( + listen: impl ToSocketAddrs, + target: String, + identity: Option, + generator: AttestationGenerator, + verifier: AttestationVerifier, + remote_certificate: Option>, + startup_check: bool, + options: TunnelOptions, + ) -> Result { + let config = tls::client_config(identity.as_ref(), remote_certificate, false)?; + Self::new_with_tls_config( + listen, + target, + config, + generator, + verifier, + identity.map(|i| i.cert_chain), + startup_check, + options, + ) + .await + } + + /// Uses the supplied certificate validation and client identity settings. + /// ALPN is replaced with the tunnel protocol. `cert_chain` must match the + /// client identity in `config`, if present, for attestation session binding. + /// `startup_check` has the same behavior as in [`Self::new`]. + #[allow(clippy::too_many_arguments)] // Mirrors new with caller-supplied TLS settings. + pub async fn new_with_tls_config( + listen: impl ToSocketAddrs, + target: String, + mut config: ClientConfig, + generator: AttestationGenerator, + verifier: AttestationVerifier, + cert_chain: Option>>, + startup_check: bool, + options: TunnelOptions, + ) -> Result { + config.alpn_protocols = vec![APPLICATION_PROTOCOL.to_vec()]; + let inner = + AttestedTlsClient::new_with_tls_config(config, generator, verifier, cert_chain)?; + let tunnel = Tunnel::bind( + listen, + normalize_target(&target, Some(443))?, + Endpoint::Client(inner.clone()), + options, + ) + .await?; + if startup_check { + tokio::time::timeout(options.setup_timeout, async { + let mut stream = connect_upstream(&inner, &tunnel.target).await?; + stream + .shutdown() + .await + .map_err(io_error("close startup check"))?; + Ok::<(), TunnelError>(()) + }) + .await + .map_err(|_| TunnelError::SetupTimeout)??; + } + Ok(Self(tunnel)) + } + + /// Return local address of the underlying listener + pub fn local_addr(&self) -> io::Result { + self.0.listener.local_addr() + } + + /// Serve until shutdown resolves, then stop accepting and drain existing + /// connections for `shutdown_grace`. Individual connection failures are logged. + pub async fn serve_until(self, shutdown: impl Future) -> Result<(), TunnelError> { + self.0.serve_until(shutdown).await + } +} + +/// Accepts attested TLS connections, verifies peers, and connects each to a fixed +/// TCP target. The target can speak any byte-stream protocol. +pub struct TunnelServer(Tunnel); + +impl TunnelServer { + pub async fn new( + listen: impl ToSocketAddrs, + target: String, + identity: TlsCertAndKey, + generator: AttestationGenerator, + verifier: AttestationVerifier, + client_auth: bool, + options: TunnelOptions, + ) -> Result { + let config = tls::server_config(&identity, client_auth)?; + Self::new_with_tls_config( + listen, + target, + config, + generator, + verifier, + identity.cert_chain, + options, + ) + .await + } + + /// Uses custom TLS settings, including private client CA policies. ALPN is + /// replaced with the tunnel protocol. `cert_chain` must match `config`. + pub async fn new_with_tls_config( + listen: impl ToSocketAddrs, + target: String, + mut config: ServerConfig, + generator: AttestationGenerator, + verifier: AttestationVerifier, + cert_chain: Vec>, + options: TunnelOptions, + ) -> Result { + config.alpn_protocols = vec![APPLICATION_PROTOCOL.to_vec()]; + let inner = + AttestedTlsServer::new_with_tls_config(cert_chain, config, generator, verifier)?; + Ok(Self( + Tunnel::bind( + listen, + normalize_target(&target, None)?, + Endpoint::Server(inner), + options, + ) + .await?, + )) + } + + /// Return local address of the underlying listener + pub fn local_addr(&self) -> io::Result { + self.0.listener.local_addr() + } + + /// Serve until shutdown resolves, then stop accepting and drain existing + /// connections for `shutdown_grace`. Individual connection failures are logged. + pub async fn serve_until(self, shutdown: impl Future) -> Result<(), TunnelError> { + self.0.serve_until(shutdown).await + } +} + +/// Attested TLS client or server +#[derive(Clone)] +enum Endpoint { + Client(AttestedTlsClient), + Server(AttestedTlsServer), +} + +impl Endpoint { + /// Common method for attested TLS connection setup for both client and server + async fn setup( + &self, + inbound: TcpStream, + target: &str, + ) -> Result<(TcpStream, tokio_rustls::TlsStream), TunnelError> { + // Disable Nagle's algorithm to reduce latency + inbound + .set_nodelay(true) + .map_err(io_error("configure inbound TCP"))?; + + match self { + Self::Client(client) => { + let stream = connect_upstream(client, target).await?; + Ok((inbound, stream.into())) + } + Self::Server(server) => { + // Do TLS handshake and attestation exchange on inbound connection + let (stream, _, _) = server.handle_connection(inbound).await?; + + // Ensure correct negotiated application protocol + require_tunnel_protocol(stream.get_ref().1.alpn_protocol())?; + + // Open connection to target service + let target = TcpStream::connect(target) + .await + .map_err(io_error("connect target"))?; + + target + .set_nodelay(true) + .map_err(io_error("configure target TCP"))?; + + Ok((target, stream.into())) + } + } + } +} + +/// Shared by startup checks and real tunnels so both enforce the same policy. +async fn connect_upstream( + client: &AttestedTlsClient, + target: &str, +) -> Result, TunnelError> { + let outbound = TcpStream::connect(target) + .await + .map_err(io_error("connect tunnel server"))?; + + outbound + .set_nodelay(true) + .map_err(io_error("configure outbound TCP"))?; + + // Do TLS handshake and attestation exchange + let (stream, _, _) = client.connect(target, outbound).await?; + + // Check application protocol was negotiated + require_tunnel_protocol(stream.get_ref().1.alpn_protocol())?; + + Ok(stream) +} + +/// Check negotiated application protocol +fn require_tunnel_protocol(protocol: Option<&[u8]>) -> Result<(), TunnelError> { + if SUPPORTED_ALPN_PROTOCOL_VERSIONS + .iter() + .any(|version| protocol == Some([*version, b"+", APPLICATION_PROTOCOL].concat().as_slice())) + { + Ok(()) + } else { + Err(TunnelError::ProtocolMismatch) + } +} + +struct Tunnel { + /// Listener for source client (for client) or proxy client (for server) + listener: TcpListener, + /// The proxy-server address (for client) or target server address (for server) + target: String, + /// Attested TLS client or server + endpoint: Endpoint, + options: TunnelOptions, +} + +impl Tunnel { + /// Setup listener and check configuration + async fn bind( + listen: impl ToSocketAddrs, + target: String, + endpoint: Endpoint, + options: TunnelOptions, + ) -> Result { + let options = options.validate()?; + + let listener = TcpListener::bind(listen) + .await + .map_err(io_error("bind listener"))?; + + Ok(Self { + listener, + target, + endpoint, + options, + }) + } + + /// Run until told to shut down + async fn serve_until(self, shutdown: impl Future) -> Result<(), TunnelError> { + let Self { + listener, + target, + endpoint, + options, + } = self; + + let slots = Arc::new(Semaphore::new(options.max_connections.get())); + + // JoinSet aborts all children when this serving future is dropped. + let mut tasks = JoinSet::new(); + + let mut next_accept = tokio::time::Instant::now(); + + tokio::pin!(shutdown); + + loop { + tokio::select! { + biased; + _ = &mut shutdown => break, + result = tasks.join_next(), if !tasks.is_empty() => { + if let Some(Err(error)) = result { + tracing::warn!(%error, "Tunnel task failed"); + } + } + incoming = async { + tokio::time::sleep_until(next_accept).await; + listener.accept().await + } => { + let (inbound, peer) = match incoming { + Ok(connection) => connection, + Err(error) => { + // Resource exhaustion and per-connection errors must + // not tear down established tunnels. Delay only this + // branch so shutdown and task reaping remain responsive. + tracing::warn!(%error, "Accept failed; retrying"); + next_accept = tokio::time::Instant::now() + Duration::from_secs(1); + continue; + } + }; + + let Ok(permit) = slots.clone().try_acquire_owned() else { + tracing::warn!(%peer, "Connection limit reached; closing new connection"); + continue; + }; + + let endpoint = endpoint.clone(); + let target = target.clone(); + let span = tracing::info_span!("tunnel", %peer, %target); + tasks.spawn(async move { + let _permit = permit; + let result = async { + let (mut local, mut remote) = tokio::time::timeout( + options.setup_timeout, endpoint.setup(inbound, &target), + ).await.map_err(|_| TunnelError::SetupTimeout)??; + tracing::debug!("Tunnel established"); + + let (sent, received) = tokio::io::copy_bidirectional(&mut local, &mut remote) + .await.map_err(io_error("forwarding"))?; + + tracing::debug!(sent, received, "Tunnel closed"); + Ok::<(), TunnelError>(()) + }.await; + + if let Err(error) = result { + tracing::warn!(%error, "Tunnel connection failed"); + } + }.instrument(span)); + } + } + } + + drop(listener); + + tracing::info!(connections = tasks.len(), "Draining tunnels"); + if tokio::time::timeout(options.shutdown_grace, async { + while let Some(result) = tasks.join_next().await { + if let Err(error) = result { + tracing::warn!(%error, "Tunnel task failed while draining"); + } + } + }) + .await + .is_err() + { + tasks.shutdown().await; + } + Ok(()) + } +} + +/// Validate and process a caller-supplied target address/hostname +fn normalize_target(target: &str, default_port: Option) -> Result { + let invalid = || { + TunnelError::Configuration( + "target must be a hostname, IPv4 address, or bracketed IPv6 address with a valid port", + ) + }; + let (host, port) = if target.starts_with('[') { + let end = target.find(']').ok_or_else(invalid)?; + target[1..end] + .parse::() + .map_err(|_| invalid())?; + let tail = &target[end + 1..]; + let port = if tail.is_empty() { + None + } else { + Some(tail.strip_prefix(':').ok_or_else(invalid)?) + }; + (&target[..=end], port) + } else { + match target.split_once(':') { + Some((host, port)) => (host, Some(port)), + None => (target, None), + } + }; + if host.is_empty() + || host + .chars() + .any(|c| c.is_whitespace() || matches!(c, '/' | '@' | '?' | '#')) + { + return Err(invalid()); + } + let port = match port { + Some(port) => port.parse::().map_err(|_| invalid())?, + None => default_port.ok_or_else(invalid)?, + }; + if port == 0 { + return Err(invalid()); + } + Ok(format!("{host}:{port}")) +} + +#[derive(Debug, thiserror::Error)] +pub enum TunnelError { + #[error("{phase}: {source}")] + Io { + phase: &'static str, + #[source] + source: io::Error, + }, + #[error("TLS/attestation: {0}")] + Attestation(#[from] AttestedTlsError), + #[error("connection setup timed out")] + SetupTimeout, + #[error("protocol validation: peer did not negotiate an attested TCP tunnel")] + ProtocolMismatch, + #[error("configuration: {0}")] + Configuration(&'static str), +} + +fn io_error(phase: &'static str) -> impl FnOnce(io::Error) -> TunnelError { + move |source| TunnelError::Io { phase, source } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn targets_and_protocols_are_unambiguous() { + for (input, expected) in [ + ("example.com", "example.com:443"), + ("127.0.0.1:42", "127.0.0.1:42"), + ("[::1]", "[::1]:443"), + ("[::1]:42", "[::1]:42"), + ] { + assert_eq!(normalize_target(input, Some(443)).unwrap(), expected); + } + for input in [ + "", + "host:0", + "host:65536", + "host:", + "host/path", + "::1", + "[oops]:443", + "[::1]oops", + "user@host:443", + ] { + assert!(normalize_target(input, Some(443)).is_err(), "{input}"); + } + assert!(normalize_target("host", None).is_err()); + assert!(require_tunnel_protocol(Some(b"flashbots-ratls/1+tcp-tunnel")).is_ok()); + for protocol in [ + None, + Some(b"flashbots-ratls/1".as_slice()), + Some(b"flashbots-ratls/1+h2".as_slice()), + Some(b"untrusted+tcp-tunnel".as_slice()), + ] { + assert!(require_tunnel_protocol(protocol).is_err()); + } + } +} diff --git a/tcp-tunnel/src/main.rs b/tcp-tunnel/src/main.rs new file mode 100644 index 0000000..a4f30dd --- /dev/null +++ b/tcp-tunnel/src/main.rs @@ -0,0 +1,392 @@ +use std::{ + fs::File, + io::BufReader, + net::SocketAddr, + num::{NonZeroU64, NonZeroUsize}, + path::{Path, PathBuf}, + time::Duration, +}; + +use anyhow::{Context, anyhow, ensure}; +use attested_tls::{ + TlsCertAndKey, + attestation::{ + AttestationGenerator, AttestationType, AttestationVerifier, PccsMode, + measurements::MeasurementPolicy, + }, + self_signed::generate_self_signed_cert, +}; +use attested_tls_tcp_tunnel::{TunnelClient, TunnelOptions, TunnelServer, tls}; +use clap::{Args, Parser, Subcommand}; +use tokio_rustls::rustls::{self, pki_types::CertificateDer}; + +#[derive(Debug, Parser)] +#[command(version = env!("GIT_REV"), about = "Forward TCP connections through remote-attested TLS")] +struct Cli { + #[command(subcommand)] + command: Command, + /// File or URL containing the remote measurement policy + #[arg( + long, + global = true, + env = "MEASUREMENTS_FILE", + conflicts_with = "allowed_remote_attestation_type" + )] + measurements_file: Option, + /// Remote attestation type to accept when no measurements file is given + #[arg(long, global = true)] + allowed_remote_attestation_type: Option, + /// PCCS URL for DCAP verification (defaults to Intel PCS) + #[arg(long, global = true)] + pccs_url: Option, + #[arg(long, global = true)] + log_debug: bool, + #[arg(long, global = true)] + log_json: bool, + /// Write DCAP quotes to quotes/ + #[arg(long, global = true)] + log_dcap_quote: bool, + #[arg(long, global = true, env = "OVERRIDE_AZURE_OUTDATED_TCB")] + override_azure_outdated_tcb: bool, +} + +#[derive(Debug, Subcommand)] +enum Command { + /// Accept local TCP connections and tunnel each to an attested server + Client { + /// Tunnel server hostname/IP and optional port (default 443) + target_addr: String, + #[arg(short, long, default_value = "127.0.0.1:0", env = "LISTEN_ADDR")] + listen_addr: SocketAddr, + #[command(flatten)] + limits: Limits, + #[command(flatten)] + identity: Identity, + /// Local attestation type (defaults to automatic detection) + #[arg(long, env = "CLIENT_ATTESTATION_TYPE")] + client_attestation_type: Option, + /// CA certificate to trust instead of public roots + #[arg(long, conflicts_with = "allow_self_signed")] + tls_ca_certificate: Option, + /// Accept a self-signed server certificate; attestation policy still applies + #[arg(long)] + allow_self_signed: bool, + }, + /// Accept attested tunnels and forward each to a fixed TCP target + Server { + /// Target service hostname/IP and required port + target_addr: String, + #[arg(short, long, default_value = "0.0.0.0:0", env = "LISTEN_ADDR")] + listen_addr: SocketAddr, + #[command(flatten)] + limits: Limits, + #[command(flatten)] + identity: Identity, + /// Local attestation type (defaults to automatic detection) + #[arg(long, env = "SERVER_ATTESTATION_TYPE")] + server_attestation_type: Option, + /// Require a client TLS certificate authenticated against public roots + #[arg(long)] + client_auth: bool, + }, +} + +#[derive(Debug, Args)] +struct Identity { + #[arg(long, env = "TLS_PRIVATE_KEY_PATH", requires = "tls_certificate_path")] + tls_private_key_path: Option, + #[arg(long, env = "TLS_CERTIFICATE_PATH", requires = "tls_private_key_path")] + tls_certificate_path: Option, + /// Dummy attestation service URL (requires local attestation type dummy) + #[arg(long)] + dev_dummy_dcap: Option, +} + +impl Identity { + fn load(&self) -> anyhow::Result> { + match (&self.tls_certificate_path, &self.tls_private_key_path) { + (None, None) => Ok(None), + (Some(cert), Some(key)) => { + let cert_chain = load_certs(cert)?; + let key_file = File::open(key) + .with_context(|| format!("Opening private key {}", key.display()))?; + let key = rustls_pemfile::private_key(&mut BufReader::new(key_file))? + .ok_or_else(|| anyhow!("No private key found in PEM"))?; + Ok(Some(TlsCertAndKey { cert_chain, key })) + } + _ => Err(anyhow!( + "Certificate chain and private key must be provided together" + )), + } + } +} + +#[derive(Debug, Args)] +struct Limits { + /// Deadline for DNS, connections, TLS, and attestation; not stream lifetime + #[arg(long, default_value = "60")] + setup_timeout_secs: NonZeroU64, + /// Maximum connections including setup; excess arrivals are closed + #[arg(long, default_value = "256", value_parser = parse_connection_limit)] + max_connections: NonZeroUsize, + /// Time to drain connections at shutdown before closing them (0 closes immediately) + #[arg(long, default_value = "30")] + shutdown_grace_secs: u64, +} + +impl From for TunnelOptions { + fn from(value: Limits) -> Self { + Self { + setup_timeout: Duration::from_secs(value.setup_timeout_secs.get()), + max_connections: value.max_connections, + shutdown_grace: Duration::from_secs(value.shutdown_grace_secs), + } + } +} + +fn parse_connection_limit(value: &str) -> Result { + let value = value.parse::().map_err(|e| e.to_string())?; + if value.get() > tokio::sync::Semaphore::MAX_PERMITS { + return Err(format!( + "must not exceed {}", + tokio::sync::Semaphore::MAX_PERMITS + )); + } + Ok(value) +} + +fn main() -> anyhow::Result<()> { + let cli = Cli::parse(); + ensure!( + cli.measurements_file.is_some() != cli.allowed_remote_attestation_type.is_some(), + "Exactly one of --measurements-file or --allowed-remote-attestation-type must be provided" + ); + rustls::crypto::aws_lc_rs::default_provider() + .install_default() + .map_err(|_| anyhow!("Failed to install Rustls crypto provider"))?; + let level = if cli.log_debug { "debug" } else { "info" }; + let filter = tracing_subscriber::EnvFilter::new(format!( + "warn,attested_tls_tcp_tunnel={level},attested_tls={level}" + )); + let subscriber = tracing_subscriber::fmt().with_env_filter(filter); + if cli.log_json { + subscriber.json().init(); + } else { + subscriber.init(); + } + + // Dropping a #[tokio::main] runtime waits indefinitely for spawn_blocking + // quote generation. After draining sockets, do not wait for those workers. + let runtime = tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build()?; + let result = runtime.block_on(run(cli)); + runtime.shutdown_timeout(Duration::ZERO); + result +} + +async fn run(cli: Cli) -> anyhow::Result<()> { + if cli.log_dcap_quote { + tokio::fs::create_dir_all("quotes").await?; + } + let policy = match cli.measurements_file { + Some(path) => MeasurementPolicy::from_file_or_url(path).await?, + None => match cli + .allowed_remote_attestation_type + .as_deref() + .unwrap_or("") + .to_lowercase() + .as_str() + { + "tdx" => MeasurementPolicy::tdx(), + name => { + let kind: AttestationType = + serde_json::from_value(serde_json::Value::String(name.to_owned()))?; + MeasurementPolicy::single_attestation_type(kind) + } + }, + }; + let mut verifier = AttestationVerifier::builder(policy) + .with_pccs_mode(PccsMode::Lazy) + .with_dump_dcap_quotes(cli.log_dcap_quote) + .with_override_azure_outdated_tcb(cli.override_azure_outdated_tcb); + if let Some(url) = cli.pccs_url { + verifier = verifier.with_pccs_url(url); + } + let verifier = verifier.build(); + match cli.command { + Command::Client { + target_addr, + listen_addr, + limits, + identity, + client_attestation_type, + tls_ca_certificate, + allow_self_signed, + } => { + let credentials = identity.load()?; + let remote_certificate = tls_ca_certificate + .as_deref() + .map(load_certs) + .transpose()? + .and_then(|certs| certs.into_iter().next()); + let config = + tls::client_config(credentials.as_ref(), remote_certificate, allow_self_signed)?; + let generator = AttestationGenerator::new_with_detection( + client_attestation_type, + identity.dev_dummy_dcap, + )?; + let client = TunnelClient::new_with_tls_config( + listen_addr, + target_addr, + config, + generator, + verifier, + credentials.map(|c| c.cert_chain), + false, // No startup check. + limits.into(), + ) + .await?; + tracing::info!(address = %client.local_addr()?, "Tunnel client listening"); + client.serve_until(shutdown_signal()).await?; + } + Command::Server { + target_addr, + listen_addr, + limits, + identity, + server_attestation_type, + client_auth, + } => { + let credentials = match identity.load()? { + Some(credentials) => credentials, + None => { + tracing::warn!( + "No TLS certificate provided; generating self-signed certificate" + ); + generate_self_signed_cert(listen_addr.ip())? + } + }; + let generator = AttestationGenerator::new_with_detection( + server_attestation_type, + identity.dev_dummy_dcap, + )?; + let server = TunnelServer::new( + listen_addr, + target_addr, + credentials, + generator, + verifier, + client_auth, + limits.into(), + ) + .await?; + tracing::info!(address = %server.local_addr()?, "Tunnel server listening"); + server.serve_until(shutdown_signal()).await?; + } + } + Ok(()) +} + +fn load_certs(path: &Path) -> anyhow::Result>> { + let certs = rustls_pemfile::certs(&mut BufReader::new( + File::open(path).with_context(|| format!("Opening certificate file {}", path.display()))?, + )) + .collect::, _>>()?; + ensure!(!certs.is_empty(), "No certificates in {}", path.display()); + Ok(certs) +} + +async fn shutdown_signal() { + #[cfg(unix)] + { + match tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) { + Ok(mut terminate) => { + tokio::select! { + result = tokio::signal::ctrl_c() => { + if let Err(error) = result { tracing::error!(%error, "Cannot listen for Ctrl-C; shutting down"); } + } + _ = terminate.recv() => {} + } + } + Err(error) => tracing::error!(%error, "Cannot listen for SIGTERM; shutting down"), + } + } + #[cfg(not(unix))] + if let Err(error) = tokio::signal::ctrl_c().await { + tracing::error!(%error, "Cannot listen for Ctrl-C; shutting down"); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn cli_defaults_and_validation() { + let Cli { + command: + Command::Client { + listen_addr, + limits, + .. + }, + .. + } = Cli::try_parse_from([ + "tunnel", + "client", + "localhost", + "--allowed-remote-attestation-type", + "none", + ]) + .unwrap() + else { + panic!("wrong command") + }; + assert_eq!(listen_addr, "127.0.0.1:0".parse().unwrap()); + assert_eq!(limits.setup_timeout_secs.get(), 60); + assert_eq!(limits.max_connections.get(), 256); + assert_eq!(limits.shutdown_grace_secs, 30); + let Cli { + command: Command::Server { listen_addr, .. }, + .. + } = Cli::try_parse_from(["tunnel", "server", "localhost:50051"]).unwrap() + else { + panic!("wrong command") + }; + assert_eq!(listen_addr, "0.0.0.0:0".parse().unwrap()); + for args in [ + vec!["--max-connections", "0"], + vec!["--setup-timeout-secs", "0"], + vec!["--tls-private-key-path", "key.pem"], + vec!["--tls-certificate-path", "cert.pem"], + vec!["--allow-self-signed", "--tls-ca-certificate", "ca.pem"], + vec![ + "--measurements-file", + "policy.json", + "--allowed-remote-attestation-type", + "none", + ], + ] { + assert!( + Cli::try_parse_from(["tunnel", "client", "localhost"].into_iter().chain(args)) + .is_err() + ); + } + assert!( + parse_connection_limit(&(tokio::sync::Semaphore::MAX_PERMITS + 1).to_string()).is_err() + ); + } + + #[test] + fn empty_and_invalid_certificate_files_fail() { + let file = tempfile::NamedTempFile::new().unwrap(); + assert!(load_certs(file.path()).is_err()); + std::fs::write( + file.path(), + "-----BEGIN CERTIFICATE-----\nnot base64\n-----END CERTIFICATE-----\n", + ) + .unwrap(); + assert!(load_certs(file.path()).is_err()); + } +} diff --git a/tcp-tunnel/src/tls.rs b/tcp-tunnel/src/tls.rs new file mode 100644 index 0000000..4b4a9e2 --- /dev/null +++ b/tcp-tunnel/src/tls.rs @@ -0,0 +1,56 @@ +//! TLS configuration helpers for the CLI and embedding applications. +use std::sync::Arc; + +use attested_tls::{AttestedTlsError, TlsCertAndKey, self_signed::SkipServerVerification}; +use tokio_rustls::rustls::{ + self, ClientConfig, RootCertStore, ServerConfig, pki_types::CertificateDer, + server::WebPkiClientVerifier, +}; + +/// Build a TLS 1.3 client configuration, preserving client authentication even +/// when accepting a self-signed server. The custom CA replaces public roots. +pub fn client_config( + identity: Option<&TlsCertAndKey>, + remote_certificate: Option>, + allow_self_signed: bool, +) -> Result { + let builder = ClientConfig::builder_with_protocol_versions(&[&rustls::version::TLS13]); + let builder = if allow_self_signed { + builder + .dangerous() + .with_custom_certificate_verifier(SkipServerVerification::new()?) + } else { + let roots = match remote_certificate { + Some(cert) => { + let mut roots = RootCertStore::empty(); + roots.add(cert)?; + roots + } + None => RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()), + }; + builder.with_root_certificates(roots) + }; + Ok(match identity { + Some(identity) => { + builder.with_client_auth_cert(identity.cert_chain.clone(), identity.key.clone_key())? + } + None => builder.with_no_client_auth(), + }) +} + +/// Build a TLS 1.3 server configuration. As in the HTTP proxy, optional client +/// certificate authentication uses public roots. Use a custom ServerConfig +/// with TunnelServer::new_with_tls_config for private client CAs. +pub fn server_config( + identity: &TlsCertAndKey, + client_auth: bool, +) -> Result { + let builder = ServerConfig::builder_with_protocol_versions(&[&rustls::version::TLS13]); + let builder = if client_auth { + let roots = RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); + builder.with_client_cert_verifier(WebPkiClientVerifier::builder(Arc::new(roots)).build()?) + } else { + builder.with_no_client_auth() + }; + Ok(builder.with_single_cert(identity.cert_chain.clone(), identity.key.clone_key())?) +} diff --git a/tcp-tunnel/tests/cli.rs b/tcp-tunnel/tests/cli.rs new file mode 100644 index 0000000..54a6301 --- /dev/null +++ b/tcp-tunnel/tests/cli.rs @@ -0,0 +1,225 @@ +#![cfg(unix)] +mod common; + +use common::*; +use std::{net::SocketAddr, process::Stdio}; +use tokio::{ + io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader}, + net::TcpStream, + process::{Child, Command}, +}; + +fn command() -> Command { + clean_command(env!("CARGO_BIN_EXE_attested-tls-tcp-tunnel")) +} + +fn clean_command(program: &str) -> Command { + let mut command = Command::new(program); + command.kill_on_drop(true); + for variable in [ + "LISTEN_ADDR", + "MEASUREMENTS_FILE", + "TLS_PRIVATE_KEY_PATH", + "TLS_CERTIFICATE_PATH", + "CLIENT_ATTESTATION_TYPE", + "SERVER_ATTESTATION_TYPE", + "OVERRIDE_AZURE_OUTDATED_TCB", + ] { + command.env_remove(variable); + } + command +} + +async fn start(args: &[&str]) -> (Child, SocketAddr) { + let (child, address, _) = start_command(command().args(args)).await; + (child, address) +} + +async fn start_command( + command: &mut Command, +) -> (Child, SocketAddr, tokio::sync::oneshot::Receiver<()>) { + let mut child = command.stdout(Stdio::piped()).spawn().unwrap(); + let mut lines = BufReader::new(child.stdout.take().unwrap()).lines(); + while let Some(line) = lines.next_line().await.unwrap() { + let value: serde_json::Value = serde_json::from_str(&line).unwrap(); + if let Some(address) = value["fields"]["address"].as_str() { + // Keep stdout open: tracing may log again during shutdown. + let (tx, rx) = tokio::sync::oneshot::channel(); + tokio::spawn(async move { + let mut tx = Some(tx); + while let Ok(Some(line)) = lines.next_line().await { + let value: serde_json::Value = serde_json::from_str(&line).unwrap(); + if value["fields"]["message"] == "Accept failed; retrying" + && let Some(tx) = tx.take() + { + let _ = tx.send(()); + } + } + }); + return (child, address.parse().unwrap(), rx); + } + } + panic!("process exited before logging its listening address"); +} + +#[tokio::test] +async fn descriptor_exhaustion_preserves_tunnels_and_recovers() { + exercise_descriptor_exhaustion(true).await; +} + +#[tokio::test] +async fn descriptor_exhaustion_does_not_delay_shutdown() { + exercise_descriptor_exhaustion(false).await; +} + +async fn exercise_descriptor_exhaustion(recover: bool) { + bounded(async { + let target = listener().await; + // Limit only the child, leaving the test runner's descriptor limit intact. + let (mut server, server_addr, accept_error) = start_command(clean_command("sh").args([ + "-c", + "ulimit -n 64 && exec \"$@\"", + "sh", + env!("CARGO_BIN_EXE_attested-tls-tcp-tunnel"), + "server", + &target.local_addr().unwrap().to_string(), + "--listen-addr", + LOCAL, + "--server-attestation-type", + "none", + "--allowed-remote-attestation-type", + "none", + "--setup-timeout-secs", + "2", + "--shutdown-grace-secs", + "0", + "--log-json", + ])) + .await; + let (mut client, client_addr) = start(&[ + "client", + &server_addr.to_string(), + "--client-attestation-type", + "none", + "--allowed-remote-attestation-type", + "none", + "--allow-self-signed", + "--shutdown-grace-secs", + "0", + "--log-json", + ]) + .await; + let mut source = TcpStream::connect(client_addr).await.unwrap(); + let (mut backend, _) = target.accept().await.unwrap(); + + let mut stalled = Vec::new(); + for _ in 0..80 { + stalled.push(TcpStream::connect(server_addr).await.unwrap()); + } + accept_error + .await + .expect("server must report descriptor exhaustion"); + assert!(server.try_wait().unwrap().is_none()); + source.write_all(b"still alive").await.unwrap(); + let mut bytes = [0; 11]; + backend.read_exact(&mut bytes).await.unwrap(); + assert_eq!(&bytes, b"still alive"); + backend.write_all(&bytes).await.unwrap(); + source.read_exact(&mut bytes).await.unwrap(); + assert_eq!(&bytes, b"still alive"); + + if !recover { + // Less than the one-second accept backoff, with descriptors still + // exhausted: shutdown must interrupt the pending retry delay. + tokio::time::timeout( + std::time::Duration::from_millis(750), + terminate(&mut server), + ) + .await + .expect("shutdown was blocked by accept backoff"); + terminate(&mut client).await; + return; + } + + drop(stalled); + // The listener must resume acceptance after descriptors become available. + let mut recovered = TcpStream::connect(client_addr).await.unwrap(); + let (mut recovered_backend, _) = target.accept().await.unwrap(); + recovered.write_all(b"recovered").await.unwrap(); + let mut bytes = [0; 9]; + recovered_backend.read_exact(&mut bytes).await.unwrap(); + assert_eq!(&bytes, b"recovered"); + terminate(&mut server).await; + terminate(&mut client).await; + }) + .await; +} + +async fn terminate(child: &mut Child) { + let status = Command::new("kill") + .args(["-TERM", &child.id().unwrap().to_string()]) + .status() + .await + .unwrap(); + assert!(status.success()); + assert!(child.wait().await.unwrap().success()); +} + +#[tokio::test] +async fn cli_round_trip_and_sigterm_shutdown() { + bounded(async { + let target = listener().await; + let target_addr = target.local_addr().unwrap().to_string(); + let (mut server, server_addr) = start(&[ + "server", + &target_addr, + "--listen-addr", + LOCAL, + "--server-attestation-type", + "none", + "--allowed-remote-attestation-type", + "none", + "--shutdown-grace-secs", + "0", + "--log-json", + ]) + .await; + let (mut client, client_addr) = start(&[ + "client", + &server_addr.to_string(), + "--client-attestation-type", + "none", + "--allowed-remote-attestation-type", + "none", + "--allow-self-signed", + "--shutdown-grace-secs", + "0", + "--log-json", + ]) + .await; + assert!(client_addr.ip().is_loopback()); + let mut source = TcpStream::connect(client_addr).await.unwrap(); + let (mut backend, _) = target.accept().await.unwrap(); + source.write_all(b"cli smoke test").await.unwrap(); + let mut bytes = [0; 14]; + backend.read_exact(&mut bytes).await.unwrap(); + backend.write_all(&bytes).await.unwrap(); + source.read_exact(&mut bytes).await.unwrap(); + assert_eq!(&bytes, b"cli smoke test"); + terminate(&mut server).await; + assert!(matches!(source.read(&mut [0; 1]).await, Ok(0) | Err(_))); + terminate(&mut client).await; + }) + .await; +} + +#[tokio::test] +async fn cli_requires_an_explicit_verification_policy() { + let output = command() + .args(["client", "127.0.0.1:443"]) + .output() + .await + .unwrap(); + assert!(!output.status.success()); + assert!(String::from_utf8_lossy(&output.stderr).contains("Exactly one of")); +} diff --git a/tcp-tunnel/tests/common/mod.rs b/tcp-tunnel/tests/common/mod.rs new file mode 100644 index 0000000..af4d2d1 --- /dev/null +++ b/tcp-tunnel/tests/common/mod.rs @@ -0,0 +1,107 @@ +#![allow(dead_code)] +use attested_tls::{ + attestation::{AttestationGenerator, AttestationVerifier}, + self_signed::generate_self_signed_cert, +}; +use attested_tls_tcp_tunnel::{TunnelClient, TunnelError, TunnelOptions, TunnelServer}; +use std::{future::Future, net::SocketAddr, time::Duration}; +use tokio::{net::TcpListener, sync::oneshot, task::JoinHandle}; + +pub const LOCAL: &str = "127.0.0.1:0"; + +pub fn provider() { + let _ = tokio_rustls::rustls::crypto::aws_lc_rs::default_provider().install_default(); +} + +pub async fn bounded(future: impl Future) -> T { + tokio::time::timeout(Duration::from_secs(15), future) + .await + .expect("test timed out") +} + +pub async fn listener() -> TcpListener { + TcpListener::bind(LOCAL).await.unwrap() +} + +pub struct Running { + pub addr: SocketAddr, + shutdown: Option>, + task: Option>>, +} + +impl Running { + pub fn client(client: TunnelClient) -> Self { + let addr = client.local_addr().unwrap(); + let (tx, rx) = oneshot::channel(); + Self { + addr, + shutdown: Some(tx), + task: Some(tokio::spawn(client.serve_until(async { + let _ = rx.await; + }))), + } + } + + pub fn server(server: TunnelServer) -> Self { + let addr = server.local_addr().unwrap(); + let (tx, rx) = oneshot::channel(); + Self { + addr, + shutdown: Some(tx), + task: Some(tokio::spawn(server.serve_until(async { + let _ = rx.await; + }))), + } + } + + pub fn signal(&mut self) { + self.shutdown.take().unwrap().send(()).unwrap(); + } + + pub fn is_finished(&self) -> bool { + self.task.as_ref().unwrap().is_finished() + } + + pub async fn wait(mut self) { + self.task.take().unwrap().await.unwrap().unwrap(); + } +} + +impl Drop for Running { + fn drop(&mut self) { + if let Some(task) = self.task.take() { + task.abort(); + } + } +} + +pub async fn pair(target: SocketAddr, options: TunnelOptions) -> (Running, Running) { + provider(); + let identity = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + let cert = identity.cert_chain[0].clone(); + let server = TunnelServer::new( + LOCAL, + target.to_string(), + identity, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + false, + options, + ) + .await + .unwrap(); + let server = Running::server(server); + let client = TunnelClient::new( + LOCAL, + server.addr.to_string(), + None, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + Some(cert), + false, // No startup check. + options, + ) + .await + .unwrap(); + (Running::client(client), server) +} diff --git a/tcp-tunnel/tests/grpc.rs b/tcp-tunnel/tests/grpc.rs new file mode 100644 index 0000000..0fdb384 --- /dev/null +++ b/tcp-tunnel/tests/grpc.rs @@ -0,0 +1,248 @@ +//! A small gRPC wire fixture: HTTP/2 DATA contains length-prefixed messages and +//! final status is carried in trailers. No protobuf compiler or external server +//! is needed to exercise the tunnel's transport guarantees. +mod common; + +use attested_tls_tcp_tunnel::TunnelOptions; +use bytes::Bytes; +use common::*; +use h2::{ + RecvStream, SendStream, + client::{ResponseFuture, SendRequest}, +}; +use std::{future::poll_fn, time::Duration}; +use tokio::net::TcpStream; + +fn message(payload: &[u8]) -> Bytes { + let mut bytes = vec![0]; + bytes.extend_from_slice(&(payload.len() as u32).to_be_bytes()); + bytes.extend_from_slice(payload); + bytes.into() +} + +async fn send(stream: &mut SendStream, mut bytes: Bytes) -> Result<(), h2::Error> { + while !bytes.is_empty() { + stream.reserve_capacity(bytes.len()); + let capacity = match poll_fn(|cx| stream.poll_capacity(cx)).await { + Some(capacity) => capacity?, + None => return Err(poll_fn(|cx| stream.poll_reset(cx)).await?.into()), + }; + if capacity > 0 { + stream.send_data(bytes.split_to(capacity.min(bytes.len())), false)?; + } + } + Ok(()) +} + +async fn collect(mut body: RecvStream, status: &str) -> Vec { + let mut bytes = Vec::new(); + while let Some(data) = body.data().await { + let data = data.unwrap(); + body.flow_control().release_capacity(data.len()).unwrap(); + bytes.extend_from_slice(&data); + } + let trailers = body.trailers().await.unwrap().unwrap(); + assert_eq!(trailers["grpc-status"], status); + assert_eq!(trailers["x-result-bin"], "AAEC"); + bytes +} + +async fn open(sender: SendRequest, method: &str) -> (ResponseFuture, SendStream) { + let mut sender = sender.ready().await.unwrap(); + let request = http::Request::builder() + .method("POST") + .uri(format!("http://application.example/test.Service/{method}")) + .header("content-type", "application/grpc") + .header("te", "trailers") + .header("x-input-bin", "AAEC") + .header("x-repeat", "first") + .header("x-repeat", "second") + .body(()) + .unwrap(); + sender.send_request(request, false).unwrap() +} + +async fn handle(request: http::Request, mut respond: h2::server::SendResponse) { + assert_eq!( + request.uri().authority().unwrap().as_str(), + "application.example" + ); + assert_eq!(request.headers()["te"], "trailers"); + assert_eq!(request.headers()["x-input-bin"], "AAEC"); + assert_eq!(request.headers().get_all("x-repeat").iter().count(), 2); + let method = request.uri().path().rsplit('/').next().unwrap().to_owned(); + let response = http::Response::builder().header("content-type", "application/grpc"); + if method == "TrailersOnly" { + respond + .send_response(response.header("grpc-status", "7").body(()).unwrap(), true) + .unwrap(); + return; + } + let mut response = respond + .send_response(response.body(()).unwrap(), false) + .unwrap(); + if method == "Cancel" { + // Exceed the response window while the caller deliberately does not read. + // RST_STREAM must still get through and leave sibling streams usable. + let result = send(&mut response, message(&vec![42; 256 * 1024])).await; + match result { + Ok(()) => assert_eq!( + poll_fn(|cx| response.poll_reset(cx)).await.unwrap(), + h2::Reason::CANCEL + ), + Err(error) => assert_eq!(error.reason(), Some(h2::Reason::CANCEL)), + } + return; + } + let mut request = request.into_body(); + let mut input = Vec::new(); + while let Some(data) = request.data().await { + let data = data.unwrap(); + request.flow_control().release_capacity(data.len()).unwrap(); + if method == "Bidi" { + send(&mut response, data).await.unwrap(); + } else { + input.extend_from_slice(&data); + } + } + if method == "ServerStream" { + send(&mut response, message(b"first")).await.unwrap(); + send(&mut response, message(b"second")).await.unwrap(); + } else if method != "Bidi" { + send(&mut response, input.into()).await.unwrap(); + } + let mut trailers = http::HeaderMap::new(); + trailers.insert( + "grpc-status", + if method == "Error" { "13" } else { "0" }.parse().unwrap(), + ); + trailers.insert("x-result-bin", "AAEC".parse().unwrap()); + response.send_trailers(trailers).unwrap(); +} + +#[tokio::test] +async fn grpc_streaming_multiplexing_trailers_and_cancellation() { + bounded(async { + let target = listener().await; + let options = TunnelOptions { + setup_timeout: Duration::from_secs(1), + ..TunnelOptions::default() + }; + let (client, _server) = pair(target.local_addr().unwrap(), options).await; + let (cancelled_tx, cancelled_rx) = tokio::sync::oneshot::channel(); + let backend = tokio::spawn(async move { + let (socket, _) = target.accept().await.unwrap(); + // All RPCs below must travel over this single target connection. + let mut connection = h2::server::handshake(socket).await.unwrap(); + let mut tasks = tokio::task::JoinSet::new(); + let mut cancelled_tx = Some(cancelled_tx); + while let Some(result) = connection.accept().await { + let (request, response) = result.unwrap(); + let cancelled = if request.uri().path().ends_with("/Cancel") { + cancelled_tx.take() + } else { + None + }; + tasks.spawn(async move { + handle(request, response).await; + if let Some(tx) = cancelled { + let _ = tx.send(()); + } + }); + } + while let Some(result) = tasks.join_next().await { + result.unwrap(); + } + }); + let socket = TcpStream::connect(client.addr).await.unwrap(); + // Leave connection-level capacity for siblings when one stream consumes + // its entire (default 64 KiB) receive window without being read. + let (sender, connection) = h2::client::Builder::new() + .initial_connection_window_size(1024 * 1024) + .handshake(socket) + .await + .unwrap(); + let driver = tokio::spawn(connection); + + let (response, mut bidi) = open(sender.clone(), "Bidi").await; + let mut bidi_response = response.await.unwrap().into_body(); + // Neither an unfinished upload nor an idle response should inherit the + // setup deadline once the tunnel has been established. + tokio::time::sleep(Duration::from_millis(1100)).await; + + let (cancelled_response, mut cancelled_upload) = open(sender.clone(), "Cancel").await; + cancelled_upload.send_data(Bytes::new(), true).unwrap(); + let held_response = cancelled_response.await.unwrap(); + + let mut calls = tokio::task::JoinSet::new(); + for method in ["Unary", "ClientStream", "ServerStream", "Error"] { + let sender = sender.clone(); + calls.spawn(async move { + let (response, mut upload) = open(sender, method).await; + let mut expected = message(b"hello").to_vec(); + // Deliberately split a gRPC message across DATA frames. + upload + .send_data(Bytes::copy_from_slice(&expected[..2]), false) + .unwrap(); + upload + .send_data(Bytes::copy_from_slice(&expected[2..]), false) + .unwrap(); + if method == "ClientStream" { + upload.send_data(message(b"again"), false).unwrap(); + expected.extend_from_slice(&message(b"again")); + } + upload.send_data(Bytes::new(), true).unwrap(); + let response = response.await.unwrap(); + assert_eq!(response.status(), 200); + assert_eq!(response.headers()["content-type"], "application/grpc"); + let body = collect( + response.into_body(), + if method == "Error" { "13" } else { "0" }, + ) + .await; + if method == "ServerStream" { + expected = [message(b"first"), message(b"second")].concat(); + } + assert_eq!(body, expected); + }); + } + for payload in [b"one".as_slice(), b"two", b"three"] { + let expected = message(payload); + bidi.send_data(expected.clone(), false).unwrap(); + let data = bidi_response.data().await.unwrap().unwrap(); + bidi_response + .flow_control() + .release_capacity(data.len()) + .unwrap(); + assert_eq!(data, expected); + } + cancelled_upload.send_reset(h2::Reason::CANCEL); + cancelled_rx.await.unwrap(); + drop(held_response); + while let Some(result) = calls.join_next().await { + result.unwrap(); + } + bidi.send_data(Bytes::new(), true).unwrap(); + assert!(collect(bidi_response, "0").await.is_empty()); + + let (response, mut upload) = open(sender.clone(), "TrailersOnly").await; + upload.send_data(Bytes::new(), true).unwrap(); + let response = response.await.unwrap(); + assert_eq!(response.headers()["grpc-status"], "7"); + assert!(response.into_body().is_end_stream()); + + // A new call after cancellation must still succeed on this connection. + let (response, mut upload) = open(sender.clone(), "Unary").await; + upload + .send_data(message(b"after cancellation"), true) + .unwrap(); + assert_eq!( + collect(response.await.unwrap().into_body(), "0").await, + message(b"after cancellation") + ); + drop(sender); + driver.abort(); + backend.abort(); + }) + .await; +} diff --git a/tcp-tunnel/tests/tunnel.rs b/tcp-tunnel/tests/tunnel.rs new file mode 100644 index 0000000..1404215 --- /dev/null +++ b/tcp-tunnel/tests/tunnel.rs @@ -0,0 +1,589 @@ +mod common; + +use attested_tls::{ + AttestedTlsClient, AttestedTlsServer, + attestation::{AttestationGenerator, AttestationType, AttestationVerifier}, + self_signed::generate_self_signed_cert, +}; +use attested_tls_tcp_tunnel::{TunnelClient, TunnelOptions, TunnelServer, tls}; +use common::*; +use std::{num::NonZeroUsize, sync::Arc, time::Duration}; +use tokio::{ + io::{AsyncReadExt, AsyncWriteExt}, + net::TcpStream, +}; +use tokio_rustls::rustls::{self, RootCertStore, server::WebPkiClientVerifier}; + +async fn closed(stream: &mut (impl tokio::io::AsyncRead + Unpin)) { + match stream.read(&mut [0; 1]).await { + Ok(0) | Err(_) => {} + result => panic!("expected EOF or transport error, got {result:?}"), + } +} + +#[tokio::test] +async fn large_binary_transfer_and_source_half_close() { + bounded(async { + let target = listener().await; + let (client, _server) = pair(target.local_addr().unwrap(), TunnelOptions::default()).await; + let payload: Vec = (0..1024 * 1024).map(|n| (n % 251) as u8).collect(); + let expected = payload.clone(); + let backend = tokio::spawn(async move { + let (mut stream, _) = target.accept().await.unwrap(); + tokio::time::sleep(Duration::from_millis(50)).await; + let mut input = Vec::new(); + stream.read_to_end(&mut input).await.unwrap(); + assert_eq!(input, expected); + input.reverse(); + stream.write_all(&input).await.unwrap(); + stream.shutdown().await.unwrap(); + }); + let mut source = TcpStream::connect(client.addr).await.unwrap(); + source.write_all(&payload).await.unwrap(); + source.shutdown().await.unwrap(); + let mut response = Vec::new(); + source.read_to_end(&mut response).await.unwrap(); + assert_eq!(response, payload.into_iter().rev().collect::>()); + backend.await.unwrap(); + }) + .await; +} + +#[tokio::test] +async fn target_half_close_preserves_remaining_upload() { + bounded(async { + let target = listener().await; + let (client, _server) = pair(target.local_addr().unwrap(), TunnelOptions::default()).await; + let backend = tokio::spawn(async move { + let (mut stream, _) = target.accept().await.unwrap(); + stream.write_all(b"greeting").await.unwrap(); + stream.shutdown().await.unwrap(); + let mut input = Vec::new(); + stream.read_to_end(&mut input).await.unwrap(); + assert_eq!(input, b"after EOF"); + }); + let mut source = TcpStream::connect(client.addr).await.unwrap(); + let mut response = Vec::new(); + source.read_to_end(&mut response).await.unwrap(); + assert_eq!(response, b"greeting"); + source.write_all(b"after EOF").await.unwrap(); + source.shutdown().await.unwrap(); + backend.await.unwrap(); + }) + .await; +} + +#[tokio::test] +async fn simultaneous_transfers_use_independent_target_connections() { + bounded(async { + let target = listener().await; + let (client, _server) = pair(target.local_addr().unwrap(), TunnelOptions::default()).await; + let backend = tokio::spawn(async move { + let mut tasks = tokio::task::JoinSet::new(); + for _ in 0..4 { + let (stream, _) = target.accept().await.unwrap(); + tasks.spawn(async move { + let (mut read, mut write) = stream.into_split(); + let output = vec![42; 512 * 1024]; + let upload = async { + let mut input = Vec::new(); + read.read_to_end(&mut input).await.unwrap(); + assert_eq!(input.len(), output.len()); + }; + let download = async { + write.write_all(&output).await.unwrap(); + write.shutdown().await.unwrap(); + }; + tokio::join!(upload, download); + }); + } + while let Some(result) = tasks.join_next().await { + result.unwrap(); + } + }); + let mut tasks = tokio::task::JoinSet::new(); + for n in 0..4 { + let addr = client.addr; + tasks.spawn(async move { + let stream = TcpStream::connect(addr).await.unwrap(); + let (mut read, mut write) = stream.into_split(); + let upload = async { + write.write_all(&vec![n; 512 * 1024]).await.unwrap(); + write.shutdown().await.unwrap(); + }; + let download = async { + let mut output = Vec::new(); + read.read_to_end(&mut output).await.unwrap(); + assert_eq!(output, vec![42; 512 * 1024]); + }; + tokio::join!(upload, download); + }); + } + while let Some(result) = tasks.join_next().await { + result.unwrap(); + } + backend.await.unwrap(); + }) + .await; +} + +#[tokio::test] +async fn client_capacity_includes_setup_and_timeout_does_not_retry() { + bounded(async { + provider(); + let stalled_server = listener().await; + let options = TunnelOptions { + setup_timeout: Duration::from_millis(150), + max_connections: NonZeroUsize::new(1).unwrap(), + ..TunnelOptions::default() + }; + let client = TunnelClient::new( + LOCAL, + stalled_server.local_addr().unwrap().to_string(), + None, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + None, + false, // No startup check. + options, + ) + .await + .unwrap(); + // Construction has no upstream side effects. + assert!( + tokio::time::timeout(Duration::from_millis(30), stalled_server.accept()) + .await + .is_err() + ); + let client = Running::client(client); + let mut first = TcpStream::connect(client.addr).await.unwrap(); + let (_stalled, _) = stalled_server.accept().await.unwrap(); + let mut excess = TcpStream::connect(client.addr).await.unwrap(); + closed(&mut excess).await; + closed(&mut first).await; + assert!( + tokio::time::timeout(Duration::from_millis(200), stalled_server.accept()) + .await + .is_err() + ); + let mut next = TcpStream::connect(client.addr).await.unwrap(); + let (_next, _) = stalled_server.accept().await.unwrap(); + closed(&mut next).await; + }) + .await; +} + +#[tokio::test] +async fn server_capacity_includes_handshakes_and_recovers() { + bounded(async { + provider(); + let target = listener().await; + let identity = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + let options = TunnelOptions { + setup_timeout: Duration::from_millis(200), + max_connections: NonZeroUsize::new(1).unwrap(), + ..TunnelOptions::default() + }; + let server = Running::server( + TunnelServer::new( + LOCAL, + target.local_addr().unwrap().to_string(), + identity, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + false, + options, + ) + .await + .unwrap(), + ); + let mut first = TcpStream::connect(server.addr).await.unwrap(); + // Send part of a TLS record, keeping the first handshake pending. + first.write_all(&[22, 3]).await.unwrap(); + tokio::time::sleep(Duration::from_millis(20)).await; + let mut excess = TcpStream::connect(server.addr).await.unwrap(); + closed(&mut excess).await; + closed(&mut first).await; + let mut next = TcpStream::connect(server.addr).await.unwrap(); + assert!( + tokio::time::timeout(Duration::from_millis(30), next.read(&mut [0; 1])) + .await + .is_err() + ); + closed(&mut next).await; + assert!( + tokio::time::timeout(Duration::from_millis(30), target.accept()) + .await + .is_err() + ); + }) + .await; +} + +#[tokio::test] +async fn established_capacity_recovers_and_graceful_shutdown_drains() { + bounded(async { + let target = listener().await; + let options = TunnelOptions { + max_connections: NonZeroUsize::new(1).unwrap(), + shutdown_grace: Duration::from_secs(2), + ..TunnelOptions::default() + }; + let (mut client, _server) = pair(target.local_addr().unwrap(), options).await; + let mut first = TcpStream::connect(client.addr).await.unwrap(); + let (mut backend, _) = target.accept().await.unwrap(); + let mut excess = TcpStream::connect(client.addr).await.unwrap(); + closed(&mut excess).await; + first.shutdown().await.unwrap(); + closed(&mut backend).await; + backend.shutdown().await.unwrap(); + closed(&mut first).await; + tokio::time::sleep(Duration::from_millis(20)).await; + let mut next = TcpStream::connect(client.addr).await.unwrap(); + let (mut backend, _) = target.accept().await.unwrap(); + client.signal(); + tokio::time::sleep(Duration::from_millis(30)).await; + assert!(!client.is_finished()); + assert!(TcpStream::connect(client.addr).await.is_err()); + next.write_all(b"still active").await.unwrap(); + let mut data = [0; 12]; + backend.read_exact(&mut data).await.unwrap(); + assert_eq!(&data, b"still active"); + next.shutdown().await.unwrap(); + backend.shutdown().await.unwrap(); + client.wait().await; + }) + .await; +} + +#[tokio::test] +async fn forced_shutdown_and_dropped_server_close_active_connections() { + bounded(async { + for abort_server in [false, true] { + let target = listener().await; + let options = TunnelOptions { + shutdown_grace: Duration::from_millis(40), + ..TunnelOptions::default() + }; + let (mut client, server) = pair(target.local_addr().unwrap(), options).await; + let mut source = TcpStream::connect(client.addr).await.unwrap(); + let (mut backend, _) = target.accept().await.unwrap(); + if abort_server { + drop(server); + } else { + client.signal(); + client.wait().await; + } + closed(&mut source).await; + closed(&mut backend).await; + } + }) + .await; +} + +#[tokio::test] +async fn refused_target_closes_source_without_retry() { + bounded(async { + let target = listener().await; + let addr = target.local_addr().unwrap(); + drop(target); + let (client, _server) = pair(addr, TunnelOptions::default()).await; + let mut source = TcpStream::connect(client.addr).await.unwrap(); + closed(&mut source).await; + let rebound = tokio::net::TcpListener::bind(addr).await.unwrap(); + assert!( + tokio::time::timeout(Duration::from_millis(150), rebound.accept()) + .await + .is_err() + ); + }) + .await; +} + +#[tokio::test] +async fn server_rejects_wrong_protocol_and_attestation_before_target_connect() { + bounded(async { + provider(); + for wrong_protocol in [true, false] { + let target = listener().await; + let identity = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + let cert = identity.cert_chain[0].clone(); + let verifier = if wrong_protocol { + AttestationVerifier::expect_none() + } else { + AttestationVerifier::mock() + }; + let server = Running::server( + TunnelServer::new( + LOCAL, + target.local_addr().unwrap().to_string(), + identity, + AttestationGenerator::with_no_attestation(), + verifier, + false, + TunnelOptions::default(), + ) + .await + .unwrap(), + ); + let mut config = tls::client_config(None, Some(cert), false).unwrap(); + config.alpn_protocols = vec![if wrong_protocol { + b"h2".to_vec() + } else { + b"tcp-tunnel".to_vec() + }]; + let client = AttestedTlsClient::new_with_tls_config( + config, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + None, + ) + .unwrap(); + let (mut stream, _, _) = client.connect_tcp(&server.addr.to_string()).await.unwrap(); + closed(&mut stream).await; + assert!( + tokio::time::timeout(Duration::from_millis(50), target.accept()) + .await + .is_err() + ); + } + }) + .await; +} + +#[tokio::test] +async fn client_rejects_wrong_protocol_attestation_and_untrusted_tls() { + bounded(async { + provider(); + for (mode, startup_check) in (0..3).flat_map(|mode| [(mode, false), (mode, true)]) { + let identity = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + let cert = identity.cert_chain[0].clone(); + let mut config = tls::server_config(&identity, false).unwrap(); + config.alpn_protocols = vec![if mode == 0 { + b"h2".to_vec() + } else { + b"tcp-tunnel".to_vec() + }]; + let server = AttestedTlsServer::new_with_tls_config( + identity.cert_chain, + config, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + ) + .unwrap(); + let incoming = listener().await; + let addr = incoming.local_addr().unwrap(); + let backend = tokio::spawn(async move { + let (socket, _) = incoming.accept().await.unwrap(); + if let Ok((mut stream, _, _)) = server.handle_connection(socket).await { + // No source payload may reach an incompatible peer. + closed(&mut stream).await; + } + }); + let verifier = if mode == 1 { + AttestationVerifier::mock() + } else { + AttestationVerifier::expect_none() + }; + let result = TunnelClient::new( + LOCAL, + addr.to_string(), + None, + AttestationGenerator::with_no_attestation(), + verifier, + if mode == 2 { None } else { Some(cert) }, + startup_check, + TunnelOptions::default(), + ) + .await; + if startup_check { + assert!( + result.is_err(), + "startup check accepted invalid peer, mode {mode}" + ); + backend.await.unwrap(); + continue; + } + let client = Running::client(result.unwrap()); + let mut source = TcpStream::connect(client.addr).await.unwrap(); + source.write_all(b"must not reach peer").await.unwrap(); + closed(&mut source).await; + backend.await.unwrap(); + } + }) + .await; +} + +#[tokio::test] +async fn self_signed_mode_preserves_client_identity_and_mutual_attestation() { + bounded(async { + provider(); + let server_identity = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + let client_identity = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + let mut roots = RootCertStore::empty(); + roots.add(client_identity.cert_chain[0].clone()).unwrap(); + let config = + rustls::ServerConfig::builder_with_protocol_versions(&[&rustls::version::TLS13]) + .with_client_cert_verifier( + WebPkiClientVerifier::builder(Arc::new(roots)) + .build() + .unwrap(), + ) + .with_single_cert(server_identity.cert_chain.clone(), server_identity.key) + .unwrap(); + let target = listener().await; + let server = Running::server( + TunnelServer::new_with_tls_config( + LOCAL, + target.local_addr().unwrap().to_string(), + config, + AttestationGenerator::new(AttestationType::DcapTdx, None).unwrap(), + AttestationVerifier::mock(), + server_identity.cert_chain, + TunnelOptions::default(), + ) + .await + .unwrap(), + ); + let config = tls::client_config(Some(&client_identity), None, true).unwrap(); + let client = Running::client( + TunnelClient::new_with_tls_config( + LOCAL, + server.addr.to_string(), + config, + AttestationGenerator::new(AttestationType::DcapTdx, None).unwrap(), + AttestationVerifier::mock(), + Some(client_identity.cert_chain), + false, // No startup check. + TunnelOptions::default(), + ) + .await + .unwrap(), + ); + let mut source = TcpStream::connect(client.addr).await.unwrap(); + let (mut backend, _) = target.accept().await.unwrap(); + source.write_all(b"authenticated").await.unwrap(); + let mut bytes = [0; 13]; + backend.read_exact(&mut bytes).await.unwrap(); + assert_eq!(&bytes, b"authenticated"); + }) + .await; +} + +#[tokio::test] +async fn application_tls_is_opaque_to_the_tunnel() { + bounded(async { + provider(); + let identity = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + let mut server_config = tls::server_config(&identity, false).unwrap(); + server_config.alpn_protocols = vec![b"h2".to_vec()]; + let mut client_config = + tls::client_config(None, Some(identity.cert_chain[0].clone()), false).unwrap(); + client_config.alpn_protocols = vec![b"h2".to_vec()]; + let target = listener().await; + let (client, _server) = pair(target.local_addr().unwrap(), TunnelOptions::default()).await; + let backend = tokio::spawn(async move { + let (socket, _) = target.accept().await.unwrap(); + let mut stream = tokio_rustls::TlsAcceptor::from(Arc::new(server_config)) + .accept(socket) + .await + .unwrap(); + assert_eq!(stream.get_ref().1.alpn_protocol(), Some(b"h2".as_slice())); + let mut data = [0; 6]; + stream.read_exact(&mut data).await.unwrap(); + stream.write_all(&data).await.unwrap(); + stream.shutdown().await.unwrap(); + }); + let socket = TcpStream::connect(client.addr).await.unwrap(); + let mut stream = tokio_rustls::TlsConnector::from(Arc::new(client_config)) + .connect( + rustls::pki_types::ServerName::try_from("127.0.0.1").unwrap(), + socket, + ) + .await + .unwrap(); + stream.write_all(b"opaque").await.unwrap(); + let mut data = Vec::new(); + stream.read_to_end(&mut data).await.unwrap(); + assert_eq!(data, b"opaque"); + stream.shutdown().await.unwrap(); + backend.await.unwrap(); + }) + .await; +} + +#[tokio::test] +async fn startup_check_closes_probe_and_serves_fresh_connections() { + bounded(async { + provider(); + let identity = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + let config = tls::client_config(None, Some(identity.cert_chain[0].clone()), false).unwrap(); + let target = listener().await; + let server = Running::server( + TunnelServer::new( + LOCAL, + target.local_addr().unwrap().to_string(), + identity, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + false, + TunnelOptions::default(), + ) + .await + .unwrap(), + ); + let client = TunnelClient::new_with_tls_config( + LOCAL, + server.addr.to_string(), + config, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + None, + true, + TunnelOptions::default(), + ) + .await + .unwrap(); + let (mut probe, _) = target.accept().await.unwrap(); + closed(&mut probe).await; + let client = Running::client(client); + let mut source = TcpStream::connect(client.addr).await.unwrap(); + let (mut backend, _) = target.accept().await.unwrap(); + source.write_all(b"fresh").await.unwrap(); + let mut bytes = [0; 5]; + backend.read_exact(&mut bytes).await.unwrap(); + assert_eq!(&bytes, b"fresh"); + }) + .await; +} + +#[tokio::test] +async fn startup_check_timeout_releases_listener() { + bounded(async { + provider(); + let stalled = listener().await; + let local = listener().await; + let address = local.local_addr().unwrap(); + drop(local); + let result = TunnelClient::new( + address, + stalled.local_addr().unwrap().to_string(), + None, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + None, + true, + TunnelOptions { + setup_timeout: Duration::from_millis(100), + ..TunnelOptions::default() + }, + ) + .await; + assert!(matches!( + result, + Err(attested_tls_tcp_tunnel::TunnelError::SetupTimeout) + )); + let _rebound = tokio::net::TcpListener::bind(address).await.unwrap(); + }) + .await; +}