diff --git a/Cargo.toml b/Cargo.toml index 3e46b57..fe88f0d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -4,18 +4,23 @@ default-members = ["crates/attested-tls-proxy"] resolver = "3" [workspace.dependencies] +anyhow = "1.0.100" 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 = "1.7.0" hyper-util = "0.1.17" 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" x509-parser = "0.18.0" diff --git a/README.md b/README.md index 218cc8b..3ee62bf 100644 --- a/README.md +++ b/README.md @@ -9,13 +9,19 @@ Details of the remote-attested TLS protocol are in [crates/attested-tls/README.m The proxy-client, on starting, immediately connects to the proxy-server and an attestation-verification exchange is made. This attested-TLS channel is then re-used for requests from that proxy-client instance. If the channel is lost, the client reconnects automatically and repeats the attestation exchange before forwarding subsequent requests. -It has five subcommands: +It has seven subcommands: - `attested-tls-proxy server` - run a proxy server, which accepts TLS connections from a proxy client, sends an attestation and then forwards traffic to a target CVM service. - `attested-tls-proxy client` - run a proxy client, which accepts connections from elsewhere, connects to and verifies the attestation from the proxy server, and then forwards traffic to it over TLS. - `attested-tls-proxy get-tls-cert` - connect to a proxy server, verify its attestation, and, if successful, write its PEM-encoded TLS certificate chain to standard output. This can be used to make subsequent connections to services using this certificate over regular TLS. - `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. +- `attested-tls-proxy tcp-tunnel-client` - forward local TCP connections through an attested tunnel. +- `attested-tls-proxy tcp-tunnel-server` - accept attested tunnels and forward each to a fixed TCP target. + +For opaque TCP forwarding, see [the TCP tunnel](crates/attested-tls-proxy/TCP_TUNNEL.md). The `tcp-tunnel-client` and `tcp-tunnel-server` commands and the `attested_tls_proxy::tcp_tunnel` library module provide one attested connection per source TCP connection, including support for gRPC. + +All commands log at info level to stderr by default. Use `--log-debug` for debug logs and `--log-json` for structured logs. Command output, such as HTTP response bodies and certificates, is written to stdout. ### How it works diff --git a/crates/attested-tls-proxy/Cargo.toml b/crates/attested-tls-proxy/Cargo.toml index 5c189ce..b3351ab 100644 --- a/crates/attested-tls-proxy/Cargo.toml +++ b/crates/attested-tls-proxy/Cargo.toml @@ -14,12 +14,12 @@ tokio = { workspace = true, features = ["full"] } tokio-rustls = { workspace = true, features = ["aws_lc_rs"] } x509-parser = { workspace = true, features = ["verify"] } thiserror.workspace = true -clap = { version = "4.5.51", features = ["derive", "env"] } -rustls-pemfile = "2.2.0" -anyhow = "1.0.100" +clap.workspace = true +rustls-pemfile.workspace = true +anyhow.workspace = true pem-rfc7468 = { version = "0.7.0", features = ["std"] } hyper = { workspace = true, features = ["server", "http2"] } -h2 = "0.4.12" +h2.workspace = true hyper-util = { workspace = true, features = ["tokio"] } http-body-util.workspace = true bytes.workspace = true @@ -31,7 +31,7 @@ reqwest = { version = "0.13.4", default-features = false, features = [ ] } webpki-roots.workspace = true tracing.workspace = true -tracing-subscriber = { version = "0.3.20", features = ["env-filter", "json"] } +tracing-subscriber.workspace = true axum = "0.8.8" tower-http = { version = "0.6.7", features = ["fs"] } rcgen.workspace = true diff --git a/crates/attested-tls-proxy/TCP_TUNNEL.md b/crates/attested-tls-proxy/TCP_TUNNEL.md new file mode 100644 index 0000000..cde4abb --- /dev/null +++ b/crates/attested-tls-proxy/TCP_TUNNEL.md @@ -0,0 +1,187 @@ +# TCP tunnel + +An attested-TLS TCP tunnel in the `attested_tls_proxy::tcp_tunnel` module, +with `tcp-tunnel-client` and `tcp-tunnel-server` commands in `attested-tls-proxy`. +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-proxy +``` + +The same binary, Docker image, and Debian package provide both HTTP proxy and TCP tunnel commands. + +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-proxy 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-proxy 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; numeric scope IDs are supported, e.g. `[fe80::1%3]:50051`. | +| `--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/`. Logs go to stderr; the default level is info. 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_proxy::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 shared `attested_tls_proxy::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-proxy --all-targets +cargo test -p attested-tls-proxy --test tcp_tunnel +``` + +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/crates/attested-tls-proxy/src/cli/attestation.rs b/crates/attested-tls-proxy/src/cli/attestation.rs new file mode 100644 index 0000000..3c04671 --- /dev/null +++ b/crates/attested-tls-proxy/src/cli/attestation.rs @@ -0,0 +1,48 @@ +use anyhow::anyhow; +use attested_tls::attestation::{ + AttestationType, AttestationVerifier, PccsMode, measurements::MeasurementPolicy, +}; + +pub(super) async fn build_verifier( + measurements_file: Option, + allowed_remote_attestation_type: Option, + pccs_url: Option, + log_dcap_quote: bool, + override_azure_outdated_tcb: bool, +) -> anyhow::Result { + if log_dcap_quote { + tokio::fs::create_dir_all("quotes").await?; + } + + let measurement_policy = match measurements_file { + Some(server_measurements) => { + MeasurementPolicy::from_file_or_url(server_measurements).await? + } + None => { + match allowed_remote_attestation_type + .ok_or(anyhow!( + "Either a measurements file or an allowed attestation type must be provided" + ))? + .to_lowercase() + .as_str() + { + "tdx" => MeasurementPolicy::tdx(), + attestation_type => { + let allowed_server_attestation_type: AttestationType = serde_json::from_value( + serde_json::Value::String(attestation_type.to_string()), + )?; + MeasurementPolicy::single_attestation_type(allowed_server_attestation_type) + } + } + } + }; + + let mut attestation_verifier_builder = AttestationVerifier::builder(measurement_policy) + .with_pccs_mode(PccsMode::Lazy) + .with_dump_dcap_quotes(log_dcap_quote) + .with_override_azure_outdated_tcb(override_azure_outdated_tcb); + if let Some(pccs_url) = pccs_url { + attestation_verifier_builder = attestation_verifier_builder.with_pccs_url(pccs_url); + } + Ok(attestation_verifier_builder.build()) +} diff --git a/crates/attested-tls-proxy/src/cli/http.rs b/crates/attested-tls-proxy/src/cli/http.rs new file mode 100644 index 0000000..f26fec1 --- /dev/null +++ b/crates/attested-tls-proxy/src/cli/http.rs @@ -0,0 +1,449 @@ +use super::pem::{ + certs_to_pem_string, load_certs_pem, load_tls_cert_and_key, load_tls_cert_and_key_server, +}; +use anyhow::{anyhow, ensure}; +use attested_tls::attestation::{ + AttestationType, AttestationVerifier, measurements::MultiMeasurements, +}; +use attested_tls_proxy::{ + AttestationGenerator, ProxyClient, ProxyClientOptions, ProxyServer, + attested_get::{attested_get, split_target_and_path}, + file_server::attested_file_server, + get_tls_cert, health_check, +}; +use clap::Args; +use std::{ + net::SocketAddr, + num::{NonZeroU64, NonZeroUsize}, + path::PathBuf, + time::Duration, +}; +use tokio::io::AsyncWriteExt; + +#[derive(Args, Debug, Clone)] +pub(super) struct ClientArgs { + /// Socket address to listen on + #[arg(short, long, default_value = "0.0.0.0:0", env = "LISTEN_ADDR")] + listen_addr: SocketAddr, + /// The hostname:port or ip:port of the proxy server (port defaults to 443) + target_addr: String, + /// Request deadline in seconds, including queueing and waiting for response headers + #[arg(long, default_value = "60")] + request_timeout_secs: NonZeroU64, + /// Close a source connection if its response body makes no write progress for this many seconds + #[arg(long, default_value = "60")] + response_body_idle_timeout_secs: NonZeroU64, + /// Maximum in-flight requests, including streaming responses + #[arg(long, default_value = "64", value_parser = parse_max_in_flight_requests)] + max_in_flight_requests: NonZeroUsize, + /// Type of attestation to present (dafaults to 'auto' for automatic detection) + /// If other than None, a TLS key and certicate must also be given + #[arg(long, env = "CLIENT_ATTESTATION_TYPE")] + client_attestation_type: Option, + /// The path to a PEM encoded private key for client authentication + #[arg(long, env = "TLS_PRIVATE_KEY_PATH")] + tls_private_key_path: Option, + /// The path to a PEM encoded certificate chain for client authentication + #[arg(long, env = "TLS_CERTIFICATE_PATH")] + tls_certificate_path: Option, + /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. + #[arg(long)] + tls_ca_certificate: Option, + /// URL of the remote dummy attestation service. Only use with --client-attestation-type + /// dummy + #[arg(long)] + dev_dummy_dcap: Option, + // Address to listen on for health checks + #[arg(long)] + listen_addr_healthcheck: Option, + /// Enables verification of self-signed TLS certificates + #[arg(long)] + allow_self_signed: bool, +} + +impl ClientArgs { + pub(super) async fn run(self, attestation_verifier: AttestationVerifier) -> anyhow::Result<()> { + let Self { + listen_addr, + target_addr, + request_timeout_secs, + response_body_idle_timeout_secs, + max_in_flight_requests, + client_attestation_type, + tls_private_key_path, + tls_certificate_path, + tls_ca_certificate, + dev_dummy_dcap, + listen_addr_healthcheck, + allow_self_signed, + } = self; + + let target_addr = target_addr + .strip_prefix("https://") + .unwrap_or(&target_addr) + .to_string(); + + if let Some(listen_addr_healthcheck) = listen_addr_healthcheck { + health_check::server(listen_addr_healthcheck).await?; + } + + let tls_cert_and_chain = if let Some(private_key) = tls_private_key_path { + Some(load_tls_cert_and_key( + tls_certificate_path + .ok_or(anyhow!("Private key given but no certificate chain"))?, + private_key, + )?) + } else { + ensure!( + tls_certificate_path.is_none(), + "Certificate chain given but no private key" + ); + None + }; + + let remote_tls_cert = match tls_ca_certificate { + Some(remote_cert_filename) => Some( + load_certs_pem(remote_cert_filename)? + .first() + .ok_or(anyhow!("Filename given but no ceritificates found"))? + .clone(), + ), + None => None, + }; + + let client_attestation_generator = + AttestationGenerator::new_with_detection(client_attestation_type, dev_dummy_dcap)?; + + let client_tls_config = attested_tls_proxy::tls::client_config( + tls_cert_and_chain.as_ref(), + remote_tls_cert, + allow_self_signed, + )?; + let client = ProxyClient::new_with_tls_config( + client_tls_config, + listen_addr, + target_addr, + client_attestation_generator, + attestation_verifier, + tls_cert_and_chain.map(|identity| identity.cert_chain), + ) + .await? + .with_request_options(ProxyClientOptions { + request_timeout: Duration::from_secs(request_timeout_secs.get()), + response_body_idle_timeout: Duration::from_secs(response_body_idle_timeout_secs.get()), + max_in_flight_requests, + }); + + loop { + if let Err(err) = client.accept().await { + tracing::error!("Failed to handle connection: {err}"); + } + } + } +} + +#[derive(Args, Debug, Clone)] +pub(super) struct ServerArgs { + /// Socket address to listen on + #[arg(short, long, default_value = "0.0.0.0:0", env = "LISTEN_ADDR")] + listen_addr: SocketAddr, + /// The hostname:port or ip:port of the target service to forward traffic to + target_addr: String, + /// Type of attestation to present (dafaults to 'auto' for automatic detection) + /// If other than None, a TLS key and certicate must also be given + #[arg(long, env = "SERVER_ATTESTATION_TYPE")] + server_attestation_type: Option, + /// The path to a PEM encoded private key + #[arg(long, env = "TLS_PRIVATE_KEY_PATH")] + tls_private_key_path: Option, + /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. + #[arg(long, env = "TLS_CERTIFICATE_PATH")] + tls_certificate_path: Option, + /// Whether to use client authentication. If the client is running in a CVM this must be + /// enabled. + #[arg(long)] + client_auth: bool, + /// URL of the remote dummy attestation service. Only use with --server-attestation-type + /// dummy + #[arg(long)] + dev_dummy_dcap: Option, + // Address to listen on for health checks + #[arg(long)] + listen_addr_healthcheck: Option, +} + +impl ServerArgs { + pub(super) async fn run(self, attestation_verifier: AttestationVerifier) -> anyhow::Result<()> { + let Self { + listen_addr, + target_addr, + tls_private_key_path, + tls_certificate_path, + client_auth, + server_attestation_type, + dev_dummy_dcap, + listen_addr_healthcheck, + } = self; + + if let Some(listen_addr_healthcheck) = listen_addr_healthcheck { + health_check::server(listen_addr_healthcheck).await?; + } + + let tls_cert_and_chain = load_tls_cert_and_key_server( + tls_certificate_path, + tls_private_key_path, + listen_addr.ip(), + )?; + + let local_attestation_generator = + AttestationGenerator::new_with_detection(server_attestation_type, dev_dummy_dcap)?; + + let server = ProxyServer::new( + tls_cert_and_chain, + listen_addr, + target_addr, + local_attestation_generator, + attestation_verifier, + client_auth, + ) + .await?; + + loop { + if let Err(err) = server.accept().await { + tracing::error!("Failed to handle connection: {err}"); + } + } + } +} + +#[derive(Args, Debug, Clone)] +pub(super) struct GetTlsCertArgs { + /// The hostname:port or ip:port of the proxy server (port defaults to 443) + server: String, + /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. + #[arg(long)] + tls_ca_certificate: Option, + /// Enables verification of self-signed TLS certificates + #[arg(long)] + allow_self_signed: bool, + /// Filename to write measurements as JSON to + #[arg(long)] + out_measurements: Option, +} + +impl GetTlsCertArgs { + pub(super) async fn run(self, attestation_verifier: AttestationVerifier) -> anyhow::Result<()> { + let Self { + server, + tls_ca_certificate, + allow_self_signed, + out_measurements, + } = self; + + let remote_tls_cert = match tls_ca_certificate { + Some(remote_cert_filename) => Some( + load_certs_pem(remote_cert_filename)? + .first() + .ok_or(anyhow!("Filename given but no ceritificates found"))? + .clone(), + ), + None => None, + }; + let (cert_chain, measurements) = get_tls_cert( + server, + attestation_verifier, + remote_tls_cert, + allow_self_signed, + ) + .await?; + + // If the user chose to write measurements to a file as JSON + if let Some(path_to_write_measurements) = out_measurements { + std::fs::write( + path_to_write_measurements, + measurements + .unwrap_or(MultiMeasurements::NoAttestation) + .to_header_format()? + .as_bytes(), + )?; + } + println!("{}", certs_to_pem_string(&cert_chain)?); + Ok(()) + } +} + +#[derive(Args, Debug, Clone)] +pub(super) struct AttestedFileServerArgs { + /// Filesystem path to statically serve + path_to_serve: PathBuf, + /// Socket address to listen on + #[arg(short, long, default_value = "0.0.0.0:0", env = "LISTEN_ADDR")] + listen_addr: SocketAddr, + /// Type of attestation to present (dafaults to none) + /// If other than None, a TLS key and certicate must also be given + #[arg(long, env = "SERVER_ATTESTATION_TYPE")] + server_attestation_type: Option, + /// The path to a PEM encoded private key + #[arg(long, env = "TLS_PRIVATE_KEY_PATH")] + tls_private_key_path: PathBuf, + /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. + #[arg(long, env = "TLS_CERTIFICATE_PATH")] + tls_certificate_path: PathBuf, + /// URL of the remote dummy attestation service. Only use with --server-attestation-type + /// dummy + #[arg(long)] + dev_dummy_dcap: Option, +} + +impl AttestedFileServerArgs { + pub(super) async fn run(self, attestation_verifier: AttestationVerifier) -> anyhow::Result<()> { + let Self { + path_to_serve, + listen_addr, + server_attestation_type, + tls_private_key_path, + tls_certificate_path, + dev_dummy_dcap, + } = self; + + let tls_cert_and_chain = load_tls_cert_and_key(tls_certificate_path, tls_private_key_path)?; + + let server_attestation_type: AttestationType = serde_json::from_value( + serde_json::Value::String(server_attestation_type.unwrap_or("none".to_string())), + )?; + + let attestation_generator = + AttestationGenerator::new(server_attestation_type, dev_dummy_dcap)?; + + attested_file_server( + path_to_serve, + tls_cert_and_chain, + listen_addr, + attestation_generator, + attestation_verifier, + false, + ) + .await?; + Ok(()) + } +} + +#[derive(Args, Debug, Clone)] +pub(super) struct AttestedGetArgs { + /// The hostname:port or ip:port of the proxy server (port defaults to 443) together + /// with the URL path to GET from the target service, eg: 127.0.0.1:3000/foobar + target_addr: String, + /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. + #[arg(long)] + tls_ca_certificate: Option, + /// Enables verification of self-signed TLS certificates + #[arg(long)] + allow_self_signed: bool, + /// Optional path to GET (defaults to '/') - this takes precedence over giving the path + /// as part of the target address. + #[arg(long)] + url_path: Option, +} + +impl AttestedGetArgs { + pub(super) async fn run(self, attestation_verifier: AttestationVerifier) -> anyhow::Result<()> { + let Self { + target_addr, + url_path, + tls_ca_certificate, + allow_self_signed, + } = self; + + let remote_tls_cert = match tls_ca_certificate { + Some(remote_cert_filename) => Some( + load_certs_pem(remote_cert_filename)? + .first() + .ok_or(anyhow!("Filename given but no ceritificates found"))? + .clone(), + ), + None => None, + }; + + let (target_addr, embedded_url_path) = split_target_and_path(&target_addr); + let url_path = url_path.or(embedded_url_path); + + let mut response = attested_get( + target_addr, + url_path.as_deref().unwrap_or("/"), + attestation_verifier, + remote_tls_cert, + allow_self_signed, + ) + .await?; + + ensure!( + !response.status().is_redirection(), + "Attested GET returned {}; redirects are not followed because the destination has not been attested", + response.status() + ); + + // Write response body to standard output + let mut stdout = tokio::io::stdout(); + + while let Some(chunk) = response.chunk().await? { + stdout.write_all(&chunk).await?; + } + + stdout.flush().await?; + Ok(()) + } +} + +/// Parses an admission limit within Tokio's supported semaphore range. +fn parse_max_in_flight_requests(value: &str) -> Result { + let count = value + .parse::() + .map_err(|error| error.to_string())?; + if count.get() > tokio::sync::Semaphore::MAX_PERMITS { + return Err(format!( + "must not exceed {}", + tokio::sync::Semaphore::MAX_PERMITS, + )); + } + Ok(count) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::cli::{Cli, CliCommand}; + use clap::Parser; + /// Checks that CLI parsing rejects invalid limits before starting the proxy. + #[test] + fn max_in_flight_requests_validates_semaphore_limit() { + for count in [ + 0, + 1, + tokio::sync::Semaphore::MAX_PERMITS, + tokio::sync::Semaphore::MAX_PERMITS + 1, + ] { + let result = Cli::try_parse_from([ + "attested-tls-proxy", + "client", + "localhost:443", + "--max-in-flight-requests", + &count.to_string(), + ]); + if (1..=tokio::sync::Semaphore::MAX_PERMITS).contains(&count) { + let CliCommand::Client(ClientArgs { + max_in_flight_requests, + .. + }) = result.unwrap().command + else { + panic!("expected client command"); + }; + assert_eq!(max_in_flight_requests.get(), count); + } else { + assert_eq!( + result.unwrap_err().kind(), + clap::error::ErrorKind::ValueValidation + ); + } + } + } +} diff --git a/crates/attested-tls-proxy/src/cli/mod.rs b/crates/attested-tls-proxy/src/cli/mod.rs new file mode 100644 index 0000000..2e39fd4 --- /dev/null +++ b/crates/attested-tls-proxy/src/cli/mod.rs @@ -0,0 +1,90 @@ +mod attestation; +mod http; +mod pem; +mod tcp_tunnel; + +use anyhow::ensure; +use clap::{Parser, Subcommand}; + +const GIT_REV: &str = match option_env!("GIT_REV") { + Some(rev) => rev, + None => "unknown", +}; + +#[derive(Parser, Debug, Clone)] +#[command(version = GIT_REV, about, long_about = None)] +pub(crate) struct Cli { + #[clap(subcommand)] + command: CliCommand, + /// Path to file, or URL, containing JSON measurements to be enforced on the remote party + #[arg(long, global = true, env = "MEASUREMENTS_FILE")] + measurements_file: Option, + /// If no measurements file is specified, a single attestion type to allow + #[arg(long, global = true)] + allowed_remote_attestation_type: Option, + /// The URL of a PCCS to use when verifying DCAP attestations. Defaults to Intel PCS. + #[arg(long, global = true)] + pccs_url: Option, + /// Log debug messages + #[arg(long, global = true)] + pub(crate) log_debug: bool, + /// Log in JSON format + #[arg(long, global = true)] + pub(crate) log_json: bool, + /// Log DCAP quotes to folder `quotes/` + #[arg(long, global = true)] + log_dcap_quote: bool, + /// Overrides Azure outdated TCB info + #[arg(long, global = true, env = "OVERRIDE_AZURE_OUTDATED_TCB")] + override_azure_outdated_tcb: bool, +} + +#[derive(Subcommand, Debug, Clone)] +enum CliCommand { + /// Accept local TCP connections and tunnel each to an attested server + TcpTunnelClient(tcp_tunnel::ClientArgs), + /// Accept attested tunnels and forward each to a fixed TCP target + TcpTunnelServer(tcp_tunnel::ServerArgs), + /// Run a proxy client + Client(http::ClientArgs), + /// Run a proxy server + Server(http::ServerArgs), + /// Retrieve the attested TLS certificate from a proxy server + GetTlsCert(http::GetTlsCertArgs), + /// Serve a filesystem path over an attested channel + AttestedFileServer(http::AttestedFileServerArgs), + /// Start a proxy-client, send a single HTTP GET request to the given path and print the + /// response to standard output + AttestedGet(http::AttestedGetArgs), +} + +impl Cli { + pub(crate) fn validate(&self) -> anyhow::Result<()> { + ensure!( + self.allowed_remote_attestation_type.is_some() != self.measurements_file.is_some(), + "Exactly one of --measurements-file or --allowed-remote-attestation-type must be provided" + ); + + Ok(()) + } + + pub(crate) async fn run(self) -> anyhow::Result<()> { + let verifier = attestation::build_verifier( + self.measurements_file, + self.allowed_remote_attestation_type, + self.pccs_url, + self.log_dcap_quote, + self.override_azure_outdated_tcb, + ) + .await?; + match self.command { + CliCommand::TcpTunnelClient(args) => args.run(verifier).await, + CliCommand::TcpTunnelServer(args) => args.run(verifier).await, + CliCommand::Client(args) => args.run(verifier).await, + CliCommand::Server(args) => args.run(verifier).await, + CliCommand::GetTlsCert(args) => args.run(verifier).await, + CliCommand::AttestedFileServer(args) => args.run(verifier).await, + CliCommand::AttestedGet(args) => args.run(verifier).await, + } + } +} diff --git a/crates/attested-tls-proxy/src/cli/pem.rs b/crates/attested-tls-proxy/src/cli/pem.rs new file mode 100644 index 0000000..274d6c2 --- /dev/null +++ b/crates/attested-tls-proxy/src/cli/pem.rs @@ -0,0 +1,186 @@ +use anyhow::anyhow; +use attested_tls::TlsCertAndKey; +use std::{fs::File, net::IpAddr, path::PathBuf}; +use tokio_rustls::rustls::pki_types::{CertificateDer, PrivateKeyDer}; + +pub(super) fn load_tls_cert_and_key_server( + cert_chain: Option, + private_key: Option, + ip: IpAddr, +) -> anyhow::Result { + if let Some(private_key) = private_key { + load_tls_cert_and_key( + cert_chain.ok_or(anyhow!("Private key given but no certificate chain"))?, + private_key, + ) + } else { + if cert_chain.is_some() { + return Err(anyhow!("Certificate chain provided but no private key")); + } + tracing::warn!("No TLS ceritifcate provided - generating self-signed"); + Ok(attested_tls_proxy::self_signed::generate_self_signed_cert( + ip, + )?) + } +} + +/// Load TLS details from storage +pub(super) fn load_tls_cert_and_key( + cert_chain: PathBuf, + private_key: PathBuf, +) -> anyhow::Result { + let key = load_private_key_pem(private_key)?; + let cert_chain = load_certs_pem(cert_chain)?; + Ok(TlsCertAndKey { key, cert_chain }) +} + +/// load certificates from a PEM-encoded file +pub(super) fn load_certs_pem(path: PathBuf) -> std::io::Result>> { + rustls_pemfile::certs(&mut std::io::BufReader::new(File::open(path)?)) + .collect::, _>>() +} + +/// load TLS private key from a PEM-encoded file +pub(super) fn load_private_key_pem(path: PathBuf) -> anyhow::Result> { + 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 +pub(super) fn certs_to_pem_string( + certs: &[CertificateDer<'_>], +) -> Result { + let mut out = String::new(); + for cert in certs { + let block = + pem_rfc7468::encode_string("CERTIFICATE", pem_rfc7468::LineEnding::LF, cert.as_ref())?; + out.push_str(&block); + out.push('\n'); + } + Ok(out) +} + +#[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(_) + )); + } +} diff --git a/crates/attested-tls-proxy/src/cli/tcp_tunnel.rs b/crates/attested-tls-proxy/src/cli/tcp_tunnel.rs new file mode 100644 index 0000000..c12d108 --- /dev/null +++ b/crates/attested-tls-proxy/src/cli/tcp_tunnel.rs @@ -0,0 +1,295 @@ +use anyhow::{anyhow, ensure}; +use attested_tls::{ + TlsCertAndKey, + attestation::{AttestationGenerator, AttestationVerifier}, +}; +use attested_tls_proxy::self_signed::generate_self_signed_cert; +use attested_tls_proxy::tcp_tunnel::{TunnelClient, TunnelOptions, TunnelServer}; +use attested_tls_proxy::tls; +use clap::Args; +use std::{ + net::SocketAddr, + num::{NonZeroU64, NonZeroUsize}, + path::PathBuf, + time::Duration, +}; +use tokio_rustls::rustls::pki_types::CertificateDer; + +#[derive(Debug, Clone, Args)] +pub(super) struct ClientArgs { + /// 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, +} + +#[derive(Debug, Clone, Args)] +pub(super) struct ServerArgs { + /// 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, Clone, 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 = super::pem::load_private_key_pem(key.clone())?; + Ok(Some(TlsCertAndKey { cert_chain, key })) + } + _ => Err(anyhow!( + "Certificate chain and private key must be provided together" + )), + } + } +} + +#[derive(Debug, Clone, 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) +} + +impl ClientArgs { + pub(super) async fn run(self, verifier: AttestationVerifier) -> anyhow::Result<()> { + let Self { + target_addr, + listen_addr, + limits, + identity, + client_attestation_type, + tls_ca_certificate, + allow_self_signed, + } = self; + + 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?; + Ok(()) + } +} + +impl ServerArgs { + pub(super) async fn run(self, verifier: AttestationVerifier) -> anyhow::Result<()> { + let Self { + target_addr, + listen_addr, + limits, + identity, + server_attestation_type, + client_auth, + } = self; + + 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: &std::path::Path) -> anyhow::Result>> { + let certs = super::pem::load_certs_pem(path.to_owned())?; + 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::*; + use crate::cli::{Cli, CliCommand}; + use clap::Parser; + + #[test] + fn cli_defaults_and_validation() { + let Cli { + command: + CliCommand::TcpTunnelClient(ClientArgs { + listen_addr, + limits, + .. + }), + .. + } = Cli::try_parse_from([ + "tunnel", + "tcp-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: CliCommand::TcpTunnelServer(ServerArgs { listen_addr, .. }), + .. + } = Cli::try_parse_from(["tunnel", "tcp-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"], + ] { + assert!( + Cli::try_parse_from( + ["tunnel", "tcp-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/crates/attested-tls-proxy/src/http/attested_get.rs b/crates/attested-tls-proxy/src/http/attested_get.rs index ba2ddd9..46e864d 100644 --- a/crates/attested-tls-proxy/src/http/attested_get.rs +++ b/crates/attested-tls-proxy/src/http/attested_get.rs @@ -32,28 +32,16 @@ pub async fn attested_get( remote_certificate: Option>, allow_self_signed: bool, ) -> Result { - let proxy_client = if allow_self_signed { - let client_config = crate::self_signed::client_tls_config_allow_self_signed(None)?; - ProxyClient::new_with_tls_config( - client_config, - "127.0.0.1:0".to_string(), - target_addr, - AttestationGenerator::with_no_attestation(), - attestation_verifier, - None, - ) - .await? - } else { - ProxyClient::new( - None, - "127.0.0.1:0".to_string(), - target_addr, - AttestationGenerator::with_no_attestation(), - attestation_verifier, - remote_certificate, - ) - .await? - }; + let client_config = crate::tls::client_config(None, remote_certificate, allow_self_signed)?; + let proxy_client = ProxyClient::new_with_tls_config( + client_config, + "127.0.0.1:0".to_string(), + target_addr, + AttestationGenerator::with_no_attestation(), + attestation_verifier, + None, + ) + .await?; attested_get_with_client(proxy_client, url_path).await } diff --git a/crates/attested-tls-proxy/src/http/mod.rs b/crates/attested-tls-proxy/src/http/mod.rs index 29c4086..298dc05 100644 --- a/crates/attested-tls-proxy/src/http/mod.rs +++ b/crates/attested-tls-proxy/src/http/mod.rs @@ -1,8 +1,9 @@ //! HTTP forwarding over attested TLS. +use crate::target::{InvalidTarget, normalize_target}; pub mod attested_get; pub mod file_server; pub mod health_check; -use crate::self_signed; +use crate::tls; pub use attested_tls; pub use attested_tls::attestation; @@ -26,10 +27,8 @@ use thiserror::Error; use tokio::io; use tokio::net::{TcpListener, TcpStream, ToSocketAddrs}; use tokio::sync::{Semaphore, mpsc, oneshot}; -use tokio_rustls::rustls::server::{VerifierBuilderError, WebPkiClientVerifier}; -use tokio_rustls::rustls::{ - self, ClientConfig, RootCertStore, ServerConfig, pki_types::CertificateDer, -}; +use tokio_rustls::rustls::server::VerifierBuilderError; +use tokio_rustls::rustls::{ClientConfig, ServerConfig, pki_types::CertificateDer}; use tracing::{debug, error, warn}; use crate::http::http_version::{ALPN_H2, ALPN_HTTP11, HttpConnection, HttpSender, HttpVersion}; @@ -76,17 +75,13 @@ pub async fn get_tls_cert( remote_certificate: Option>, allow_self_signed: bool, ) -> Result<(Vec>, Option), AttestedTlsError> { - let (cert, measurements) = if allow_self_signed { - let client_tls_config = self_signed::client_tls_config_allow_self_signed(None)?; - attested_tls::get_tls_cert_with_config( - &server_name, - attestation_verifier, - client_tls_config, - ) - .await? - } else { - attested_tls::get_tls_cert(server_name, attestation_verifier, remote_certificate).await? - }; + let client_tls_config = tls::client_config(None, remote_certificate, allow_self_signed)?; + let (cert, measurements) = attested_tls::get_tls_cert_with_config( + &server_name, + attestation_verifier, + client_tls_config, + ) + .await?; debug!("[get-tls-cert] Connected to proxy server with measurements: {measurements:?}"); Ok((cert, measurements)) @@ -100,6 +95,8 @@ pub struct ProxyServer { listener: Arc, /// The address/hostname of the target service we are proxying to target: String, + /// Normalized TCP destination; preserve the original target for the Host header. + target_addr: String, } impl ProxyServer { @@ -111,25 +108,8 @@ impl ProxyServer { attestation_verifier: AttestationVerifier, client_auth: bool, ) -> Result { - let mut server_config = if client_auth { - let root_store = - RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); - let verifier = WebPkiClientVerifier::builder(Arc::new(root_store)).build()?; - - ServerConfig::builder_with_protocol_versions(&[&rustls::version::TLS13]) - .with_client_cert_verifier(verifier) - .with_single_cert( - cert_and_key.cert_chain.clone(), - cert_and_key.key.clone_key(), - )? - } else { - ServerConfig::builder_with_protocol_versions(&[&rustls::version::TLS13]) - .with_no_client_auth() - .with_single_cert( - cert_and_key.cert_chain.clone(), - cert_and_key.key.clone_key(), - )? - }; + let target_addr = normalize_target(&target, None)?; + let mut server_config = tls::server_config(&cert_and_key, client_auth)?; ensure_proxy_alpn_protocols(&mut server_config.alpn_protocols); let attested_tls_server = AttestedTlsServer::new_with_tls_config( @@ -145,6 +125,7 @@ impl ProxyServer { attested_tls_server, listener: listener.into(), target, + target_addr, }) } @@ -157,6 +138,7 @@ impl ProxyServer { attestation_generator: AttestationGenerator, attestation_verifier: AttestationVerifier, ) -> Result { + let target_addr = normalize_target(&target, None)?; ensure_proxy_alpn_protocols(&mut server_config.alpn_protocols); let attested_tls_server = AttestedTlsServer::new_with_tls_config( @@ -172,6 +154,7 @@ impl ProxyServer { attested_tls_server, listener: listener.into(), target, + target_addr, }) } @@ -180,6 +163,7 @@ impl ProxyServer { /// Returns the handle for the task handling the connection pub async fn accept(&self) -> Result, ProxyError> { let target = self.target.clone(); + let target_addr = self.target_addr.clone(); let (inbound, client_addr) = self.listener.accept().await?; let attested_tls_server = self.attested_tls_server.clone(); @@ -191,6 +175,7 @@ impl ProxyServer { measurements, attestation_type, target, + target_addr, client_addr, ) .await @@ -218,6 +203,7 @@ impl ProxyServer { measurements: Option, remote_attestation_type: AttestationType, target: String, + target_addr: String, client_addr: SocketAddr, ) -> Result<(), ProxyError> { debug!("[proxy-server] accepted connection with measurements: {measurements:?}"); @@ -270,7 +256,7 @@ impl ProxyServer { remote_attestation_type.as_str(), ); - let target = target.clone(); + let target = target_addr.clone(); async move { match Self::handle_http_request(req, target).await { Ok(res) => { @@ -379,27 +365,8 @@ impl ProxyClient { attestation_verifier: AttestationVerifier, remote_certificate: Option>, ) -> Result { - let root_store = match remote_certificate { - Some(remote_certificate) => { - let mut root_store = RootCertStore::empty(); - root_store.add(remote_certificate)?; - root_store - } - None => RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()), - }; - - let mut client_config = if let Some(ref cert_and_key) = cert_and_key { - ClientConfig::builder_with_protocol_versions(&[&rustls::version::TLS13]) - .with_root_certificates(root_store) - .with_client_auth_cert( - cert_and_key.cert_chain.clone(), - cert_and_key.key.clone_key(), - )? - } else { - ClientConfig::builder_with_protocol_versions(&[&rustls::version::TLS13]) - .with_root_certificates(root_store) - .with_no_client_auth() - }; + let mut client_config = + tls::client_config(cert_and_key.as_ref(), remote_certificate, false)?; ensure_proxy_alpn_protocols(&mut client_config.alpn_protocols); let attested_tls_client = AttestedTlsClient::new_with_tls_config( @@ -439,11 +406,9 @@ impl ProxyClient { attested_tls_client: AttestedTlsClient, target_name: &str, ) -> Result { + let target = normalize_target(target_name, Some(443))?; let listener = TcpListener::bind(address).await?; - // Process the hostname / port provided by the user - let target = host_to_host_with_port(target_name); - // Channel for getting incoming requests from the source client let (requests_tx, mut requests_rx) = mpsc::channel::(1024); @@ -763,6 +728,8 @@ where /// An error when running a proxy client or server #[derive(Error, Debug)] pub enum ProxyError { + #[error("Invalid target: {0}")] + InvalidTarget(#[from] InvalidTarget), #[error("Failed to get server ceritifcate")] NoCertificate, #[error("TLS: {0}")] @@ -799,15 +766,6 @@ impl From> for ProxyError { } } -/// If no port was provided, default to 443 -pub(crate) fn host_to_host_with_port(host: &str) -> String { - if host.contains(':') { - host.to_string() - } else { - format!("{host}:443") - } -} - /// An Executor for hyper that uses the tokio runtime #[derive(Clone)] pub(crate) struct TokioExecutor; @@ -1620,6 +1578,7 @@ mod tests { attested_tls_server, listener: listener.into(), target: target_addr.to_string(), + target_addr: target_addr.to_string(), }; let proxy_addr = proxy_server.local_addr().unwrap(); diff --git a/crates/attested-tls-proxy/src/lib.rs b/crates/attested-tls-proxy/src/lib.rs index 82a5428..20a0172 100644 --- a/crates/attested-tls-proxy/src/lib.rs +++ b/crates/attested-tls-proxy/src/lib.rs @@ -1,6 +1,10 @@ -//! An attested TLS protocol and HTTPS proxy. +//! HTTP and TCP proxies over attested TLS. pub mod http; pub mod self_signed; +mod target; +pub mod tcp_tunnel; +pub mod tls; +pub use target::InvalidTarget; // Preserve the original HTTP proxy API at the crate root. pub use http::*; diff --git a/crates/attested-tls-proxy/src/main.rs b/crates/attested-tls-proxy/src/main.rs index 0a9a7dd..dfa21ef 100644 --- a/crates/attested-tls-proxy/src/main.rs +++ b/crates/attested-tls-proxy/src/main.rs @@ -1,720 +1,44 @@ -use anyhow::{anyhow, ensure}; -use attested_tls::attestation::measurements::MultiMeasurements; -use clap::{Parser, Subcommand}; -use std::{ - fs::File, - net::{IpAddr, SocketAddr}, - num::{NonZeroU64, NonZeroUsize}, - path::PathBuf, - time::Duration, -}; -use tokio::io::AsyncWriteExt; -use tokio_rustls::rustls::pki_types::{CertificateDer, PrivateKeyDer}; -use tracing::level_filters::LevelFilter; - -use attested_tls_proxy::{ - AttestationGenerator, ProxyClient, ProxyClientOptions, ProxyServer, - attested_get::{attested_get, split_target_and_path}, - attested_tls::{ - TlsCertAndKey, - attestation::{ - AttestationType, AttestationVerifier, PccsMode, measurements::MeasurementPolicy, - }, - }, - file_server::attested_file_server, - get_tls_cert, health_check, -}; - -const GIT_REV: &str = match option_env!("GIT_REV") { - Some(rev) => rev, - None => "unknown", -}; - -#[derive(Parser, Debug, Clone)] -#[command(version = GIT_REV, about, long_about = None)] -struct Cli { - #[clap(subcommand)] - command: CliCommand, - /// Path to file, or URL, containing JSON measurements to be enforced on the remote party - #[arg(long, global = true, env = "MEASUREMENTS_FILE")] - measurements_file: Option, - /// If no measurements file is specified, a single attestion type to allow - #[arg(long, global = true)] - allowed_remote_attestation_type: Option, - /// The URL of a PCCS to use when verifying DCAP attestations. Defaults to Intel PCS. - #[arg(long, global = true)] - pccs_url: Option, - /// Log debug messages - #[arg(long, global = true)] - log_debug: bool, - /// Log in JSON format - #[arg(long, global = true)] - log_json: bool, - /// Log DCAP quotes to folder `quotes/` - #[arg(long, global = true)] - log_dcap_quote: bool, - /// Overrides Azure outdated TCB info - #[arg(long, global = true, env = "OVERRIDE_AZURE_OUTDATED_TCB")] - override_azure_outdated_tcb: bool, -} +mod cli; -#[derive(Subcommand, Debug, Clone)] -enum CliCommand { - /// Run a proxy client - Client { - /// Socket address to listen on - #[arg(short, long, default_value = "0.0.0.0:0", env = "LISTEN_ADDR")] - listen_addr: SocketAddr, - /// The hostname:port or ip:port of the proxy server (port defaults to 443) - target_addr: String, - /// Request deadline in seconds, including queueing and waiting for response headers - #[arg(long, default_value = "60")] - request_timeout_secs: NonZeroU64, - /// Close a source connection if its response body makes no write progress for this many seconds - #[arg(long, default_value = "60")] - response_body_idle_timeout_secs: NonZeroU64, - /// Maximum in-flight requests, including streaming responses - #[arg(long, default_value = "64", value_parser = parse_max_in_flight_requests)] - max_in_flight_requests: NonZeroUsize, - /// Type of attestation to present (dafaults to 'auto' for automatic detection) - /// If other than None, a TLS key and certicate must also be given - #[arg(long, env = "CLIENT_ATTESTATION_TYPE")] - client_attestation_type: Option, - /// The path to a PEM encoded private key for client authentication - #[arg(long, env = "TLS_PRIVATE_KEY_PATH")] - tls_private_key_path: Option, - /// The path to a PEM encoded certificate chain for client authentication - #[arg(long, env = "TLS_CERTIFICATE_PATH")] - tls_certificate_path: Option, - /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. - #[arg(long)] - tls_ca_certificate: Option, - /// URL of the remote dummy attestation service. Only use with --client-attestation-type - /// dummy - #[arg(long)] - dev_dummy_dcap: Option, - // Address to listen on for health checks - #[arg(long)] - listen_addr_healthcheck: Option, - /// Enables verification of self-signed TLS certificates - #[arg(long)] - allow_self_signed: bool, - }, - /// Run a proxy server - Server { - /// Socket address to listen on - #[arg(short, long, default_value = "0.0.0.0:0", env = "LISTEN_ADDR")] - listen_addr: SocketAddr, - /// The hostname:port or ip:port of the target service to forward traffic to - target_addr: String, - /// Type of attestation to present (dafaults to 'auto' for automatic detection) - /// If other than None, a TLS key and certicate must also be given - #[arg(long, env = "SERVER_ATTESTATION_TYPE")] - server_attestation_type: Option, - /// The path to a PEM encoded private key - #[arg(long, env = "TLS_PRIVATE_KEY_PATH")] - tls_private_key_path: Option, - /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. - #[arg(long, env = "TLS_CERTIFICATE_PATH")] - tls_certificate_path: Option, - /// Whether to use client authentication. If the client is running in a CVM this must be - /// enabled. - #[arg(long)] - client_auth: bool, - /// URL of the remote dummy attestation service. Only use with --server-attestation-type - /// dummy - #[arg(long)] - dev_dummy_dcap: Option, - // Address to listen on for health checks - #[arg(long)] - listen_addr_healthcheck: Option, - }, - /// Retrieve the attested TLS certificate from a proxy server - GetTlsCert { - /// The hostname:port or ip:port of the proxy server (port defaults to 443) - server: String, - /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. - #[arg(long)] - tls_ca_certificate: Option, - /// Enables verification of self-signed TLS certificates - #[arg(long)] - allow_self_signed: bool, - /// Filename to write measurements as JSON to - #[arg(long)] - out_measurements: Option, - }, - /// Serve a filesystem path over an attested channel - AttestedFileServer { - /// Filesystem path to statically serve - path_to_serve: PathBuf, - /// Socket address to listen on - #[arg(short, long, default_value = "0.0.0.0:0", env = "LISTEN_ADDR")] - listen_addr: SocketAddr, - /// Type of attestation to present (dafaults to none) - /// If other than None, a TLS key and certicate must also be given - #[arg(long, env = "SERVER_ATTESTATION_TYPE")] - server_attestation_type: Option, - /// The path to a PEM encoded private key - #[arg(long, env = "TLS_PRIVATE_KEY_PATH")] - tls_private_key_path: PathBuf, - /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. - #[arg(long, env = "TLS_CERTIFICATE_PATH")] - tls_certificate_path: PathBuf, - /// URL of the remote dummy attestation service. Only use with --server-attestation-type - /// dummy - #[arg(long)] - dev_dummy_dcap: Option, - }, - /// Start a proxy-client, send a single HTTP GET request to the given path and print the - /// response to standard output - AttestedGet { - /// The hostname:port or ip:port of the proxy server (port defaults to 443) together - /// with the URL path to GET from the target service, eg: 127.0.0.1:3000/foobar - target_addr: String, - /// Additional CA certificate to verify against (PEM) Defaults to no additional TLS certs. - #[arg(long)] - tls_ca_certificate: Option, - /// Enables verification of self-signed TLS certificates - #[arg(long)] - allow_self_signed: bool, - /// Optional path to GET (defaults to '/') - this takes precedence over giving the path - /// as part of the target address. - #[arg(long)] - url_path: Option, - }, -} +use anyhow::anyhow; +use clap::Parser; +use cli::Cli; +use std::time::Duration; +use tracing::level_filters::LevelFilter; -#[tokio::main] -async fn main() -> anyhow::Result<()> { +fn main() -> anyhow::Result<()> { + let cli = Cli::parse(); tokio_rustls::rustls::crypto::aws_lc_rs::default_provider() .install_default() .map_err(|_| anyhow!("Failed to install the default rustls crypto provider"))?; - - let cli = Cli::parse(); - - ensure!( - cli.allowed_remote_attestation_type.is_some() != cli.measurements_file.is_some(), - "Exactly one of --measurements-file or --allowed-remote-attestation-type must be provided" - ); - + cli.validate()?; + init_logging(&cli); + let runtime = tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build()?; + let result = runtime.block_on(cli.run()); + // Command-specific cleanup has finished. Do not wait indefinitely for + // blocking quote generation when tearing down the runtime. + runtime.shutdown_timeout(Duration::ZERO); + result +} + +fn init_logging(cli: &Cli) { let crate_name = env!("CARGO_CRATE_NAME"); + let level = if cli.log_debug { "debug" } else { "info" }; + let filter = format!("{crate_name}={level},attested_tls={level}"); let env_filter = tracing_subscriber::EnvFilter::builder() .with_default_directive(LevelFilter::WARN.into()) // global default - .parse_lossy(format!( - "{crate_name}={}", - if cli.log_debug { "debug" } else { "warn" } - )); + .parse_lossy(filter); - let subscriber = tracing_subscriber::fmt::Subscriber::builder().with_env_filter(env_filter); + let subscriber = tracing_subscriber::fmt::Subscriber::builder() + .with_env_filter(env_filter) + .with_writer(std::io::stderr); if cli.log_json { subscriber.json().init(); } else { subscriber.pretty().init(); } - - if cli.log_dcap_quote { - tokio::fs::create_dir_all("quotes").await?; - } - - let measurement_policy = match cli.measurements_file { - Some(server_measurements) => { - MeasurementPolicy::from_file_or_url(server_measurements).await? - } - None => { - match cli - .allowed_remote_attestation_type - .ok_or(anyhow!( - "Either a measurements file or an allowed attestation type must be provided" - ))? - .to_lowercase() - .as_str() - { - "tdx" => MeasurementPolicy::tdx(), - attestation_type => { - let allowed_server_attestation_type: AttestationType = serde_json::from_value( - serde_json::Value::String(attestation_type.to_string()), - )?; - MeasurementPolicy::single_attestation_type(allowed_server_attestation_type) - } - } - } - }; - - let mut attestation_verifier_builder = AttestationVerifier::builder(measurement_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(pccs_url) = cli.pccs_url { - attestation_verifier_builder = attestation_verifier_builder.with_pccs_url(pccs_url); - } - let attestation_verifier = attestation_verifier_builder.build(); - - match cli.command { - CliCommand::Client { - listen_addr, - target_addr, - request_timeout_secs, - response_body_idle_timeout_secs, - max_in_flight_requests, - client_attestation_type, - tls_private_key_path, - tls_certificate_path, - tls_ca_certificate, - dev_dummy_dcap, - listen_addr_healthcheck, - allow_self_signed, - } => { - let target_addr = target_addr - .strip_prefix("https://") - .unwrap_or(&target_addr) - .to_string(); - - if let Some(listen_addr_healthcheck) = listen_addr_healthcheck { - health_check::server(listen_addr_healthcheck).await?; - } - - let tls_cert_and_chain = if let Some(private_key) = tls_private_key_path { - Some(load_tls_cert_and_key( - tls_certificate_path - .ok_or(anyhow!("Private key given but no certificate chain"))?, - private_key, - )?) - } else { - ensure!( - tls_certificate_path.is_none(), - "Certificate chain given but no private key" - ); - None - }; - - let remote_tls_cert = match tls_ca_certificate { - Some(remote_cert_filename) => Some( - load_certs_pem(remote_cert_filename)? - .first() - .ok_or(anyhow!("Filename given but no ceritificates found"))? - .clone(), - ), - None => None, - }; - - let client_attestation_generator = - AttestationGenerator::new_with_detection(client_attestation_type, dev_dummy_dcap)?; - - let client = if allow_self_signed { - let client_tls_config = - attested_tls_proxy::self_signed::client_tls_config_allow_self_signed( - tls_cert_and_chain.as_ref(), - )?; - ProxyClient::new_with_tls_config( - client_tls_config, - listen_addr, - target_addr, - client_attestation_generator, - attestation_verifier, - tls_cert_and_chain.map(|identity| identity.cert_chain), - ) - .await? - } else { - ProxyClient::new( - tls_cert_and_chain, - listen_addr, - target_addr, - client_attestation_generator, - attestation_verifier, - remote_tls_cert, - ) - .await? - } - .with_request_options(ProxyClientOptions { - request_timeout: Duration::from_secs(request_timeout_secs.get()), - response_body_idle_timeout: Duration::from_secs( - response_body_idle_timeout_secs.get(), - ), - max_in_flight_requests, - }); - - loop { - if let Err(err) = client.accept().await { - tracing::error!("Failed to handle connection: {err}"); - } - } - } - CliCommand::Server { - listen_addr, - target_addr, - tls_private_key_path, - tls_certificate_path, - client_auth, - server_attestation_type, - dev_dummy_dcap, - listen_addr_healthcheck, - } => { - if let Some(listen_addr_healthcheck) = listen_addr_healthcheck { - health_check::server(listen_addr_healthcheck).await?; - } - - let tls_cert_and_chain = load_tls_cert_and_key_server( - tls_certificate_path, - tls_private_key_path, - listen_addr.ip(), - )?; - - let local_attestation_generator = - AttestationGenerator::new_with_detection(server_attestation_type, dev_dummy_dcap)?; - - let server = ProxyServer::new( - tls_cert_and_chain, - listen_addr, - target_addr, - local_attestation_generator, - attestation_verifier, - client_auth, - ) - .await?; - - loop { - if let Err(err) = server.accept().await { - tracing::error!("Failed to handle connection: {err}"); - } - } - } - CliCommand::GetTlsCert { - server, - tls_ca_certificate, - allow_self_signed, - out_measurements, - } => { - let remote_tls_cert = match tls_ca_certificate { - Some(remote_cert_filename) => Some( - load_certs_pem(remote_cert_filename)? - .first() - .ok_or(anyhow!("Filename given but no ceritificates found"))? - .clone(), - ), - None => None, - }; - let (cert_chain, measurements) = get_tls_cert( - server, - attestation_verifier, - remote_tls_cert, - allow_self_signed, - ) - .await?; - - // If the user chose to write measurements to a file as JSON - if let Some(path_to_write_measurements) = out_measurements { - std::fs::write( - path_to_write_measurements, - measurements - .unwrap_or(MultiMeasurements::NoAttestation) - .to_header_format()? - .as_bytes(), - )?; - } - println!("{}", certs_to_pem_string(&cert_chain)?); - } - CliCommand::AttestedFileServer { - path_to_serve, - listen_addr, - server_attestation_type, - tls_private_key_path, - tls_certificate_path, - dev_dummy_dcap, - } => { - let tls_cert_and_chain = - load_tls_cert_and_key(tls_certificate_path, tls_private_key_path)?; - - let server_attestation_type: AttestationType = serde_json::from_value( - serde_json::Value::String(server_attestation_type.unwrap_or("none".to_string())), - )?; - - let attestation_generator = - AttestationGenerator::new(server_attestation_type, dev_dummy_dcap)?; - - attested_file_server( - path_to_serve, - tls_cert_and_chain, - listen_addr, - attestation_generator, - attestation_verifier, - false, - ) - .await?; - } - CliCommand::AttestedGet { - target_addr, - url_path, - tls_ca_certificate, - allow_self_signed, - } => { - let remote_tls_cert = match tls_ca_certificate { - Some(remote_cert_filename) => Some( - load_certs_pem(remote_cert_filename)? - .first() - .ok_or(anyhow!("Filename given but no ceritificates found"))? - .clone(), - ), - None => None, - }; - - let (target_addr, embedded_url_path) = split_target_and_path(&target_addr); - let url_path = url_path.or(embedded_url_path); - - let mut response = attested_get( - target_addr, - url_path.as_deref().unwrap_or("/"), - attestation_verifier, - remote_tls_cert, - allow_self_signed, - ) - .await?; - - ensure!( - !response.status().is_redirection(), - "Attested GET returned {}; redirects are not followed because the destination has not been attested", - response.status() - ); - - // Write response body to standard output - let mut stdout = tokio::io::stdout(); - - while let Some(chunk) = response.chunk().await? { - stdout.write_all(&chunk).await?; - } - - stdout.flush().await?; - } - } - - Ok(()) -} - -fn load_tls_cert_and_key_server( - cert_chain: Option, - private_key: Option, - ip: IpAddr, -) -> anyhow::Result { - if let Some(private_key) = private_key { - load_tls_cert_and_key( - cert_chain.ok_or(anyhow!("Private key given but no certificate chain"))?, - private_key, - ) - } else { - if cert_chain.is_some() { - return Err(anyhow!("Certificate chain provided but no private key")); - } - tracing::warn!("No TLS ceritifcate provided - generating self-signed"); - Ok(attested_tls_proxy::self_signed::generate_self_signed_cert( - ip, - )?) - } -} - -/// Load TLS details from storage -fn load_tls_cert_and_key( - cert_chain: PathBuf, - private_key: PathBuf, -) -> anyhow::Result { - let key = load_private_key_pem(private_key)?; - let cert_chain = load_certs_pem(cert_chain)?; - Ok(TlsCertAndKey { key, cert_chain }) -} - -/// load certificates from a PEM-encoded file -fn load_certs_pem(path: PathBuf) -> std::io::Result>> { - rustls_pemfile::certs(&mut std::io::BufReader::new(File::open(path)?)) - .collect::, _>>() -} - -/// load TLS private key from a PEM-encoded file -fn load_private_key_pem(path: PathBuf) -> anyhow::Result> { - 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 -fn certs_to_pem_string(certs: &[CertificateDer<'_>]) -> Result { - let mut out = String::new(); - for cert in certs { - let block = - pem_rfc7468::encode_string("CERTIFICATE", pem_rfc7468::LineEnding::LF, cert.as_ref())?; - out.push_str(&block); - out.push('\n'); - } - Ok(out) -} - -/// Parses an admission limit within Tokio's supported semaphore range. -fn parse_max_in_flight_requests(value: &str) -> Result { - let count = value - .parse::() - .map_err(|error| error.to_string())?; - if count.get() > tokio::sync::Semaphore::MAX_PERMITS { - return Err(format!( - "must not exceed {}", - tokio::sync::Semaphore::MAX_PERMITS, - )); - } - Ok(count) -} - -#[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] - fn max_in_flight_requests_validates_semaphore_limit() { - for count in [ - 0, - 1, - tokio::sync::Semaphore::MAX_PERMITS, - tokio::sync::Semaphore::MAX_PERMITS + 1, - ] { - let result = Cli::try_parse_from([ - "attested-tls-proxy", - "client", - "localhost:443", - "--max-in-flight-requests", - &count.to_string(), - ]); - if (1..=tokio::sync::Semaphore::MAX_PERMITS).contains(&count) { - let CliCommand::Client { - max_in_flight_requests, - .. - } = result.unwrap().command - else { - panic!("expected client command"); - }; - assert_eq!(max_in_flight_requests.get(), count); - } else { - assert_eq!( - result.unwrap_err().kind(), - clap::error::ErrorKind::ValueValidation - ); - } - } - } } diff --git a/crates/attested-tls-proxy/src/self_signed.rs b/crates/attested-tls-proxy/src/self_signed.rs index 648f3a6..b4f0b15 100644 --- a/crates/attested-tls-proxy/src/self_signed.rs +++ b/crates/attested-tls-proxy/src/self_signed.rs @@ -23,21 +23,6 @@ pub fn generate_self_signed_cert(ip_address: IpAddr) -> Result, -) -> Result { - let builder = rustls::ClientConfig::builder() - .dangerous() - .with_custom_certificate_verifier(SkipServerVerification::new()?); - Ok(match identity { - Some(identity) => { - builder.with_client_auth_cert(identity.cert_chain.clone(), identity.key.clone_key())? - } - None => builder.with_no_client_auth(), - }) -} - /// Used to allow verification of self-signed certificates #[derive(Debug, Clone)] pub struct SkipServerVerification { @@ -246,7 +231,7 @@ mod tests { server.handle_connection(tcp_stream).await.unwrap(); }); - let client_config = client_tls_config_allow_self_signed(None).unwrap(); + let client_config = crate::tls::client_config(None, None, true).unwrap(); let client = AttestedTlsClient::new_with_tls_config( client_config.into(), @@ -304,7 +289,7 @@ mod tests { }); // Inner TLS config - let client_config = client_tls_config_allow_self_signed(None).unwrap(); + let client_config = crate::tls::client_config(None, None, true).unwrap(); let client = AttestedTlsClient::new_with_tls_config( client_config.into(), diff --git a/crates/attested-tls-proxy/src/target.rs b/crates/attested-tls-proxy/src/target.rs new file mode 100644 index 0000000..aed98c1 --- /dev/null +++ b/crates/attested-tls-proxy/src/target.rs @@ -0,0 +1,106 @@ +/// A target is not a host/IP with a valid, nonzero port. +#[derive(Debug, Clone, Copy, thiserror::Error)] +#[error("target must be a hostname, IPv4 address, or bracketed IPv6 address with a valid port")] +pub struct InvalidTarget; + +/// Validate and process a caller-supplied target address/hostname +pub(crate) fn normalize_target( + target: &str, + default_port: Option, +) -> Result { + let invalid = || InvalidTarget; + let (host, port) = if target.starts_with('[') { + let end = target.find(']').ok_or_else(invalid)?; + // SocketAddrV6 validates both the IPv6 address and an optional numeric + // scope ID. Keep the original host text for the actual connection. + format!("{}:0", &target[..=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}")) +} + +#[cfg(test)] +mod tests { + use super::*; + #[test] + fn targets_are_unambiguous() { + for (input, expected) in [ + ("example.com", "example.com:443"), + ("127.0.0.1:42", "127.0.0.1:42"), + ("example.com:00443", "example.com:443"), + ("[::1]", "[::1]:443"), + ("[::1]:42", "[::1]:42"), + ("[fe80::1%3]", "[fe80::1%3]:443"), + ("[fe80::1%3]:8080", "[fe80::1%3]:8080"), + ("[fe80::1%0]:8080", "[fe80::1%0]:8080"), + ] { + 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", + "host?query:443", + "host#fragment:443", + "host name:443", + "[fe80::1%]:443", + "[fe80::1%-1]:443", + "[fe80::1%4294967296]:443", + "[fe80::1%3%4]:443", + "[fe80::1%eth0]:443", + ] { + assert!(normalize_target(input, Some(443)).is_err(), "{input}"); + } + assert!(normalize_target("host", None).is_err()); + assert!(normalize_target("host", Some(0)).is_err()); + assert_eq!(normalize_target("[::1]:80", None).unwrap(), "[::1]:80"); + } + #[test] + fn scoped_ipv6_target_preserves_routing_information() { + use std::net::{SocketAddr, ToSocketAddrs}; + let target = normalize_target("[fe80::1%3]:08080", None).unwrap(); + assert_eq!(target, "[fe80::1%3]:8080"); + let SocketAddr::V6(address) = target.to_socket_addrs().unwrap().next().unwrap() else { + panic!("expected IPv6 address"); + }; + assert_eq!(address.scope_id(), 3); + assert_eq!(address.port(), 8080); + assert_eq!( + *address.ip(), + "fe80::1".parse::().unwrap() + ); + } +} diff --git a/crates/attested-tls-proxy/src/tcp_tunnel/mod.rs b/crates/attested-tls-proxy/src/tcp_tunnel/mod.rs new file mode 100644 index 0000000..030bc7a --- /dev/null +++ b/crates/attested-tls-proxy/src/tcp_tunnel/mod.rs @@ -0,0 +1,461 @@ +//! 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. +use crate::tls; + +use crate::target::{InvalidTarget, normalize_target}; + +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(()) + } +} + +#[derive(Debug, thiserror::Error)] +pub enum TunnelError { + #[error("configuration: {0}")] + InvalidTarget(#[from] InvalidTarget), + #[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 protocols_are_unambiguous() { + 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/crates/attested-tls-proxy/src/tls.rs b/crates/attested-tls-proxy/src/tls.rs new file mode 100644 index 0000000..940cad9 --- /dev/null +++ b/crates/attested-tls-proxy/src/tls.rs @@ -0,0 +1,57 @@ +//! Shared TLS 1.3 configuration for HTTP proxies and TCP tunnels. +use std::sync::Arc; + +use crate::self_signed::SkipServerVerification; +use attested_tls::{AttestedTlsError, TlsCertAndKey}; +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. Optional client certificate authentication +/// uses public roots. For private client CAs, supply a custom ServerConfig through +/// the HTTP proxy or TCP tunnel's `new_with_tls_config` constructor. +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/crates/attested-tls-proxy/tests/attested_get_redirect.rs b/crates/attested-tls-proxy/tests/http/attested_get_redirect.rs similarity index 100% rename from crates/attested-tls-proxy/tests/attested_get_redirect.rs rename to crates/attested-tls-proxy/tests/http/attested_get_redirect.rs diff --git a/crates/attested-tls-proxy/tests/http/main.rs b/crates/attested-tls-proxy/tests/http/main.rs new file mode 100644 index 0000000..f633074 --- /dev/null +++ b/crates/attested-tls-proxy/tests/http/main.rs @@ -0,0 +1,2 @@ +mod attested_get_redirect; +mod target; diff --git a/crates/attested-tls-proxy/tests/http/target.rs b/crates/attested-tls-proxy/tests/http/target.rs new file mode 100644 index 0000000..9a98503 --- /dev/null +++ b/crates/attested-tls-proxy/tests/http/target.rs @@ -0,0 +1,113 @@ +use std::time::Duration; + +use attested_tls_proxy::{ + AttestationGenerator, ProxyClient, ProxyError, ProxyServer, attestation::AttestationVerifier, + attested_get::attested_get, self_signed::generate_self_signed_cert, +}; +use axum::{Router, http::HeaderMap, routing::get}; +use tokio::{net::TcpListener, time::timeout}; + +#[tokio::test] +async fn normalized_connection_preserves_original_host_header() { + check_target("127.0.0.1:0", "127.0.0.1").await; +} + +#[tokio::test] +async fn scoped_ipv6_connection_preserves_original_host_header() { + // Scope zero works on loopback without depending on interface indices. + check_target("[::1]:0", "[::1%0]").await; +} + +async fn check_target(listen: &str, host: &str) { + let _ = tokio_rustls::rustls::crypto::aws_lc_rs::default_provider().install_default(); + let listener = TcpListener::bind(listen).await.unwrap(); + let server_ip = listener.local_addr().unwrap().ip(); + // A leading zero is removed from the TCP port but must remain in Host. + let target = format!("{host}:0{}", listener.local_addr().unwrap().port()); + let backend = tokio::spawn(async move { + let app = Router::new().route( + "/", + get(|headers: HeaderMap| async move { headers["host"].to_str().unwrap().to_owned() }), + ); + axum::serve(listener, app).await.unwrap(); + }); + let server = ProxyServer::new( + generate_self_signed_cert(server_ip).unwrap(), + listen, + target.clone(), + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + false, + ) + .await + .unwrap(); + let address = server.local_addr().unwrap(); + let proxy = tokio::spawn(async move { server.accept().await.unwrap() }); + timeout(Duration::from_secs(5), async { + let response = attested_get( + format!("{host}:{}", address.port()), + "/", + AttestationVerifier::expect_none(), + None, + true, + ) + .await + .unwrap(); + assert_eq!(response.status(), http::StatusCode::OK); + assert_eq!(response.text().await.unwrap(), target); + }) + .await + .unwrap(); + proxy.await.unwrap().abort(); + backend.abort(); +} + +#[tokio::test] +async fn invalid_http_targets_fail_at_construction() { + let _ = tokio_rustls::rustls::crypto::aws_lc_rs::default_provider().install_default(); + for target in ["", "localhost:0", "localhost:65536", "localhost/path"] { + let result = ProxyClient::new( + None, + "127.0.0.1:0", + target.into(), + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + None, + ) + .await; + assert!( + matches!(result, Err(ProxyError::InvalidTarget(_))), + "{target}" + ); + } + // Server targets require an explicit port, with either TLS constructor. + for custom_tls in [false, true] { + let identity = generate_self_signed_cert("127.0.0.1".parse().unwrap()).unwrap(); + let result = if custom_tls { + let config = tokio_rustls::rustls::ServerConfig::builder() + .with_no_client_auth() + .with_single_cert(identity.cert_chain.clone(), identity.key) + .unwrap(); + ProxyServer::new_with_tls_config( + identity.cert_chain, + config, + "127.0.0.1:0", + "localhost".into(), + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + ) + .await + } else { + ProxyServer::new( + identity, + "127.0.0.1:0", + "localhost".into(), + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + false, + ) + .await + }; + assert!(matches!(result, Err(ProxyError::InvalidTarget(_)))); + } +} diff --git a/crates/attested-tls-proxy/tests/tcp_tunnel/cli.rs b/crates/attested-tls-proxy/tests/tcp_tunnel/cli.rs new file mode 100644 index 0000000..7d0ea3b --- /dev/null +++ b/crates/attested-tls-proxy/tests/tcp_tunnel/cli.rs @@ -0,0 +1,237 @@ +#![cfg(unix)] + +use super::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-proxy")) +} + +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.stderr(Stdio::piped()).spawn().unwrap(); + let mut lines = BufReader::new(child.stderr.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 stderr 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-proxy"), + "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(&[ + "tcp-tunnel-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(&[ + "tcp-tunnel-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(&[ + "tcp-tunnel-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() { + for subcommand in ["tcp-tunnel-client", "tcp-tunnel-server"] { + for policy in [ + vec![], + vec![ + "--measurements-file", + "unused.json", + "--allowed-remote-attestation-type", + "none", + ], + ] { + let output = command() + .args([subcommand, "127.0.0.1:443"]) + .args(policy) + .output() + .await + .unwrap(); + assert!(!output.status.success()); + assert!(String::from_utf8_lossy(&output.stderr).contains("Exactly one of")); + } + } +} diff --git a/crates/attested-tls-proxy/tests/tcp_tunnel/common.rs b/crates/attested-tls-proxy/tests/tcp_tunnel/common.rs new file mode 100644 index 0000000..cd73492 --- /dev/null +++ b/crates/attested-tls-proxy/tests/tcp_tunnel/common.rs @@ -0,0 +1,105 @@ +#![allow(dead_code)] +use attested_tls::attestation::{AttestationGenerator, AttestationVerifier}; +use attested_tls_proxy::self_signed::generate_self_signed_cert; +use attested_tls_proxy::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/crates/attested-tls-proxy/tests/tcp_tunnel/grpc.rs b/crates/attested-tls-proxy/tests/tcp_tunnel/grpc.rs new file mode 100644 index 0000000..775a05b --- /dev/null +++ b/crates/attested-tls-proxy/tests/tcp_tunnel/grpc.rs @@ -0,0 +1,247 @@ +//! 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. + +use super::common::*; +use attested_tls_proxy::tcp_tunnel::TunnelOptions; +use bytes::Bytes; +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/crates/attested-tls-proxy/tests/tcp_tunnel/main.rs b/crates/attested-tls-proxy/tests/tcp_tunnel/main.rs new file mode 100644 index 0000000..e52036b --- /dev/null +++ b/crates/attested-tls-proxy/tests/tcp_tunnel/main.rs @@ -0,0 +1,4 @@ +mod cli; +mod common; +mod grpc; +mod tunnel; diff --git a/crates/attested-tls-proxy/tests/tcp_tunnel/tunnel.rs b/crates/attested-tls-proxy/tests/tcp_tunnel/tunnel.rs new file mode 100644 index 0000000..70710f8 --- /dev/null +++ b/crates/attested-tls-proxy/tests/tcp_tunnel/tunnel.rs @@ -0,0 +1,589 @@ +use attested_tls_proxy::self_signed::generate_self_signed_cert; +use attested_tls_proxy::tls; + +use super::common::*; +use attested_tls::{ + AttestedTlsClient, AttestedTlsServer, + attestation::{AttestationGenerator, AttestationType, AttestationVerifier}, +}; +use attested_tls_proxy::tcp_tunnel::{TunnelClient, TunnelOptions, TunnelServer}; +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_proxy::tcp_tunnel::TunnelError::SetupTimeout) + )); + let _rebound = tokio::net::TcpListener::bind(address).await.unwrap(); + }) + .await; +} diff --git a/crates/attested-tls/src/lib.rs b/crates/attested-tls/src/lib.rs index 7cc68c7..05e0050 100644 --- a/crates/attested-tls/src/lib.rs +++ b/crates/attested-tls/src/lib.rs @@ -580,6 +580,11 @@ where pub(crate) fn server_name_from_host( host: &str, ) -> Result, tokio_rustls::rustls::pki_types::InvalidDnsNameError> { + // A scope ID selects the local network interface, not the TLS peer identity. + if let Ok(std::net::SocketAddr::V6(address)) = host.parse() { + return ServerName::try_from(address.ip().to_string()); + } + // If host contains ':', try to split off the port. let host_part = host.rsplit_once(':').map(|(h, _)| h).unwrap_or(host); @@ -629,6 +634,16 @@ fn map_alpn_protocols(existing_protocols: Vec>) -> Vec> { #[cfg(test)] mod tests { + #[test] + fn scoped_ipv6_tls_identity_excludes_scope_id() { + for target in ["[fe80::1%3]:443", "[fe80::1%0]:8443", "[fe80::1]:443"] { + assert_eq!( + super::server_name_from_host(target).unwrap(), + tokio_rustls::rustls::pki_types::ServerName::try_from("fe80::1").unwrap(), + ); + } + } + use super::*; use crate::test_helpers::{generate_certificate_chain, generate_tls_config};