use anyhow::{anyhow, Context, Result};
use quinn::congestion::BbrConfig;
use quinn::crypto::rustls::{QuicClientConfig, QuicServerConfig};
use quinn::{ClientConfig, Endpoint};
use quinn::{IdleTimeout, ServerConfig, TransportConfig, VarInt};
use rcgen::generate_simple_self_signed;
use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, ServerName, UnixTime};
use rustls::{ClientConfig as TlsClientConfig, ServerConfig as TlsServerConfig};
use std::fs;
use std::net::{IpAddr, SocketAddr};
use std::path::Path;
use std::{sync::Arc, time::Duration};
use tracing::{debug, info, warn};
use crate::common::proxy::{create_socks5_proxied_socket, ProxyConfig};
use crate::common::tls::{cert_sha256, format_fingerprint, ClientTlsConfig, ServerTlsConfig};
static ALPN_QUIC_HTTP: &[&[u8]] = &[b"h3"];
#[derive(Debug, Clone, Copy, Default)]
pub enum Congestion {
#[default]
Cubic,
Bbr,
}
fn build_transport_config(congestion: Congestion) -> Arc<TransportConfig> {
let mut tc = TransportConfig::default();
let idle_timeout = IdleTimeout::try_from(Duration::from_secs(30))
.expect("30s fits in QUIC idle-timeout VarInt");
tc.stream_receive_window(VarInt::from_u32(16 * 1024 * 1024))
.receive_window(VarInt::from_u32(64 * 1024 * 1024))
.send_window(64 * 1024 * 1024)
.keep_alive_interval(Some(Duration::from_secs(15)))
.max_idle_timeout(Some(idle_timeout))
.max_concurrent_bidi_streams(VarInt::from_u32(1024))
.max_concurrent_uni_streams(VarInt::from_u32(0));
if let Congestion::Bbr = congestion {
tc.congestion_controller_factory(Arc::new(BbrConfig::default()));
}
Arc::new(tc)
}
pub fn create_server_endpoint(
host: IpAddr,
port: u16,
tls: &ServerTlsConfig,
congestion: Congestion,
) -> Result<Endpoint> {
let addr: SocketAddr = SocketAddr::new(host, port);
let (cert, key) = load_server_identity(tls)?;
let mut server_config = build_quic_server_config(tls, cert, key)?;
server_config.transport_config(build_transport_config(congestion));
Ok(Endpoint::server(server_config, addr)?)
}
pub fn create_client_endpoint(
tls: &ClientTlsConfig,
congestion: Congestion,
server_addr: SocketAddr,
) -> Result<Endpoint> {
let mut client_config = build_quic_client_config(tls)?;
client_config.transport_config(build_transport_config(congestion));
let bind_addr: SocketAddr = match server_addr {
SocketAddr::V4(_) => "0.0.0.0:0".parse()?,
SocketAddr::V6(_) => "[::]:0".parse()?,
};
let mut endpoint = Endpoint::client(bind_addr)?;
endpoint.set_default_client_config(client_config);
Ok(endpoint)
}
pub async fn create_client_endpoint_via_proxy(
tls: &ClientTlsConfig,
congestion: Congestion,
server_addr: SocketAddr,
proxy: &ProxyConfig,
) -> Result<Endpoint> {
use quinn::{EndpointConfig, TokioRuntime};
use std::sync::Arc;
let mut client_config = build_quic_client_config(tls)?;
client_config.transport_config(build_transport_config(congestion));
let socket = create_socks5_proxied_socket(proxy, server_addr).await?;
let mut endpoint = Endpoint::new_with_abstract_socket(
EndpointConfig::default(),
None,
socket,
Arc::new(TokioRuntime),
)?;
endpoint.set_default_client_config(client_config);
Ok(endpoint)
}
fn load_server_identity(
tls: &ServerTlsConfig,
) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>)> {
let (cert_chain, key) = match tls {
ServerTlsConfig::Insecure => {
warn!(
"starting server in --insecure mode: ephemeral self-signed cert, \
no client authentication. DO NOT use in production."
);
generate_ephemeral_self_signed()?
}
ServerTlsConfig::SelfSigned { state_dir } => load_or_create_self_signed(state_dir)?,
ServerTlsConfig::Provided { cert, key } => load_pem_identity(cert, key)?,
ServerTlsConfig::Mtls { cert, key, .. } => load_pem_identity(cert, key)?,
};
if let Some(leaf) = cert_chain.first() {
let fp = format_fingerprint(&cert_sha256(leaf));
info!(fingerprint = %fp, "server cert");
}
Ok((cert_chain, key))
}
fn load_or_create_self_signed(
state_dir: &Path,
) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>)> {
let cert_path = state_dir.join("server.pem");
let key_path = state_dir.join("server.key");
if cert_path.exists() && key_path.exists() {
debug!(dir = %state_dir.display(), "loading persisted self-signed identity");
return load_pem_identity(&cert_path, &key_path);
}
info!(dir = %state_dir.display(), "generating new self-signed identity");
fs::create_dir_all(state_dir)
.with_context(|| format!("failed to create state dir {}", state_dir.display()))?;
let generated = generate_simple_self_signed(vec!["localhost".into()])
.context("failed to generate self-signed certificate")?;
let cert_pem = generated.cert.pem();
let key_pem = generated.signing_key.serialize_pem();
fs::write(&cert_path, &cert_pem)
.with_context(|| format!("failed to write {}", cert_path.display()))?;
write_secret_file(&key_path, key_pem.as_bytes())
.with_context(|| format!("failed to write {}", key_path.display()))?;
info!(
cert = %cert_path.display(),
key = %key_path.display(),
"persisted server identity",
);
let cert_der: CertificateDer<'static> = generated.cert.into();
let key_der: PrivateKeyDer<'static> =
PrivatePkcs8KeyDer::from(generated.signing_key.serialize_der()).into();
Ok((vec![cert_der], key_der))
}
fn load_pem_identity(
cert_path: &Path,
key_path: &Path,
) -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>)> {
let cert_chain = load_pem_certs(cert_path)?;
let key_pem = fs::read(key_path)
.with_context(|| format!("failed to read key file {}", key_path.display()))?;
let mut key_reader = std::io::BufReader::new(key_pem.as_slice());
let key = rustls_pemfile::private_key(&mut key_reader)
.with_context(|| format!("failed to parse PEM private key in {}", key_path.display()))?
.ok_or_else(|| anyhow!("no private key found in {}", key_path.display()))?;
Ok((cert_chain, key))
}
fn load_pem_certs(path: &Path) -> Result<Vec<CertificateDer<'static>>> {
let pem =
fs::read(path).with_context(|| format!("failed to read cert file {}", path.display()))?;
let mut reader = std::io::BufReader::new(pem.as_slice());
let certs: Vec<CertificateDer<'static>> = rustls_pemfile::certs(&mut reader)
.collect::<std::result::Result<_, _>>()
.with_context(|| format!("failed to parse PEM certs in {}", path.display()))?;
if certs.is_empty() {
return Err(anyhow!("no certificates found in {}", path.display()));
}
Ok(certs)
}
fn load_root_store(ca_path: &Path) -> Result<rustls::RootCertStore> {
let mut roots = rustls::RootCertStore::empty();
let mut added = 0usize;
for cert in load_pem_certs(ca_path)? {
roots.add(cert).with_context(|| {
format!(
"failed to add CA cert from {} to root store",
ca_path.display()
)
})?;
added += 1;
}
debug!(count = added, path = %ca_path.display(), "loaded CA certs");
Ok(roots)
}
fn write_secret_file(path: &Path, contents: &[u8]) -> Result<()> {
fs::write(path, contents)?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
let mut perms = fs::metadata(path)?.permissions();
perms.set_mode(0o600);
fs::set_permissions(path, perms)?;
}
Ok(())
}
fn build_quic_server_config(
tls: &ServerTlsConfig,
cert: Vec<CertificateDer<'static>>,
key: PrivateKeyDer<'static>,
) -> Result<ServerConfig> {
let mut server_crypto: TlsServerConfig = match tls {
ServerTlsConfig::Insecure
| ServerTlsConfig::SelfSigned { .. }
| ServerTlsConfig::Provided { .. } => TlsServerConfig::builder()
.with_no_client_auth()
.with_single_cert(cert, key)?,
ServerTlsConfig::Mtls { ca, .. } => {
info!(ca = %ca.display(), "mTLS enabled (requiring client cert)");
let roots = load_root_store(ca)?;
let verifier = rustls::server::WebPkiClientVerifier::builder(Arc::new(roots))
.build()
.context("failed to build client cert verifier")?;
TlsServerConfig::builder()
.with_client_cert_verifier(verifier)
.with_single_cert(cert, key)?
}
};
server_crypto.alpn_protocols = ALPN_QUIC_HTTP.iter().map(|&x| x.into()).collect();
Ok(ServerConfig::with_crypto(Arc::new(
QuicServerConfig::try_from(server_crypto)?,
)))
}
fn build_quic_client_config(tls: &ClientTlsConfig) -> Result<ClientConfig> {
let mut client_crypto = match tls {
ClientTlsConfig::Insecure => {
warn!(
"starting client in --insecure mode: skipping server certificate verification. \
MITM-vulnerable; for testing only."
);
TlsClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(SkipServerVerification::new())
.with_no_client_auth()
}
ClientTlsConfig::Fingerprint { sha256, .. } => TlsClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(FingerprintVerifier::new(*sha256))
.with_no_client_auth(),
ClientTlsConfig::Ca { ca, .. } => {
let roots = load_root_store(ca)?;
TlsClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth()
}
ClientTlsConfig::Mtls { ca, cert, key, .. } => {
let roots = load_root_store(ca)?;
let (cert_chain, key) = load_pem_identity(cert, key)?;
if let Some(leaf) = cert_chain.first() {
debug!(fingerprint = %format_fingerprint(&cert_sha256(leaf)), "client cert");
}
TlsClientConfig::builder()
.with_root_certificates(roots)
.with_client_auth_cert(cert_chain, key)
.context("failed to install client auth cert")?
}
};
client_crypto.alpn_protocols = ALPN_QUIC_HTTP.iter().map(|&x| x.into()).collect();
Ok(ClientConfig::new(Arc::new(QuicClientConfig::try_from(
client_crypto,
)?)))
}
pub fn client_server_name(tls: &ClientTlsConfig, server_host: &str) -> String {
match tls {
ClientTlsConfig::Insecure => server_host.to_string(),
ClientTlsConfig::Fingerprint { server_name, .. }
| ClientTlsConfig::Ca { server_name, .. }
| ClientTlsConfig::Mtls { server_name, .. } => server_name
.clone()
.unwrap_or_else(|| server_host.to_string()),
}
}
fn generate_ephemeral_self_signed() -> Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>)>
{
debug!("generating ephemeral self-signed certificate");
let cert = generate_simple_self_signed(vec!["localhost".into()])?;
let key = PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der());
let cert = cert.cert.into();
Ok((vec![cert], key.into()))
}
#[derive(Debug)]
struct SkipServerVerification(Arc<rustls::crypto::CryptoProvider>);
impl SkipServerVerification {
fn new() -> Arc<Self> {
Arc::new(Self(Arc::new(rustls::crypto::ring::default_provider())))
}
}
impl rustls::client::danger::ServerCertVerifier for SkipServerVerification {
fn verify_server_cert(
&self,
_end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp: &[u8],
_now: UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&self.0.signature_verification_algorithms,
)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&self.0.signature_verification_algorithms,
)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
self.0.signature_verification_algorithms.supported_schemes()
}
}
#[derive(Debug)]
struct FingerprintVerifier {
expected: [u8; 32],
crypto: Arc<rustls::crypto::CryptoProvider>,
}
impl FingerprintVerifier {
fn new(expected: [u8; 32]) -> Arc<Self> {
Arc::new(Self {
expected,
crypto: Arc::new(rustls::crypto::ring::default_provider()),
})
}
}
impl rustls::client::danger::ServerCertVerifier for FingerprintVerifier {
fn verify_server_cert(
&self,
end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp: &[u8],
_now: UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
let actual = cert_sha256(end_entity);
if actual == self.expected {
Ok(rustls::client::danger::ServerCertVerified::assertion())
} else {
warn!(
expected = %format_fingerprint(&self.expected),
actual = %format_fingerprint(&actual),
"server cert fingerprint mismatch",
);
Err(rustls::Error::InvalidCertificate(
rustls::CertificateError::ApplicationVerificationFailure,
))
}
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&self.crypto.signature_verification_algorithms,
)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&self.crypto.signature_verification_algorithms,
)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
self.crypto
.signature_verification_algorithms
.supported_schemes()
}
}