use std::net::{SocketAddr, ToSocketAddrs};
use std::sync::Arc;
use std::time::Duration;
use quinn::{ClientConfig, Endpoint, IdleTimeout, TransportConfig};
use crate::cert::{PubkeyPinVerifier, SkipServerVerification};
pub const ALPN: &[u8] = b"macula";
pub const DEFAULT_IDLE_TIMEOUT: Duration = Duration::from_secs(300);
pub const DEFAULT_KEEP_ALIVE_INTERVAL: Duration = Duration::from_secs(15);
#[derive(Clone, Copy)]
pub enum Trust {
Pinned([u8; 32]),
WebPki,
Insecure,
}
#[derive(Debug)]
pub enum ConnectError {
Resolve(std::io::Error),
NoAddress,
Endpoint(std::io::Error),
Config(rustls::Error),
Connect(quinn::ConnectError),
Connection(quinn::ConnectionError),
}
impl std::fmt::Display for ConnectError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ConnectError::Resolve(e) => write!(f, "resolving station address: {e}"),
ConnectError::NoAddress => write!(f, "hostname resolved to no addresses"),
ConnectError::Endpoint(e) => write!(f, "creating QUIC endpoint: {e}"),
ConnectError::Config(e) => write!(f, "building TLS config: {e}"),
ConnectError::Connect(e) => write!(f, "starting QUIC connect: {e}"),
ConnectError::Connection(e) => write!(f, "QUIC connection failed: {e}"),
}
}
}
impl std::error::Error for ConnectError {}
fn client_config(trust: Trust) -> Result<ClientConfig, rustls::Error> {
let mut crypto = match trust {
Trust::Pinned(pubkey) => rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(PubkeyPinVerifier::new(pubkey)))
.with_no_client_auth(),
Trust::WebPki => {
let mut roots = rustls::RootCertStore::empty();
roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
rustls::ClientConfig::builder()
.with_root_certificates(roots)
.with_no_client_auth()
}
Trust::Insecure => rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(SkipServerVerification::new()))
.with_no_client_auth(),
};
crypto.alpn_protocols = vec![ALPN.to_vec()];
let mut transport = TransportConfig::default();
transport.max_idle_timeout(Some(
IdleTimeout::try_from(DEFAULT_IDLE_TIMEOUT).expect("valid idle timeout"),
));
transport.keep_alive_interval(Some(DEFAULT_KEEP_ALIVE_INTERVAL));
apply_flow_control_defaults(&mut transport);
let quic_crypto = quinn::crypto::rustls::QuicClientConfig::try_from(crypto)
.map_err(|e| rustls::Error::General(e.to_string()))?;
let mut config = ClientConfig::new(Arc::new(quic_crypto));
config.transport_config(Arc::new(transport));
Ok(config)
}
fn apply_flow_control_defaults(transport: &mut TransportConfig) {
transport.stream_receive_window((16u32 * 1024 * 1024).into());
transport.receive_window((64u32 * 1024 * 1024).into());
transport.send_window(64u64 * 1024 * 1024);
}
pub async fn connect(
host: &str,
port: u16,
trust: Trust,
) -> Result<quinn::Connection, ConnectError> {
let addr = resolve(host, port)?;
let bind_addr: SocketAddr = if addr.is_ipv6() {
"[::]:0".parse().expect("valid unspecified v6 addr")
} else {
"0.0.0.0:0".parse().expect("valid unspecified v4 addr")
};
let mut endpoint = Endpoint::client(bind_addr).map_err(ConnectError::Endpoint)?;
let config = client_config(trust).map_err(ConnectError::Config)?;
endpoint.set_default_client_config(config);
let connecting = endpoint
.connect(addr, host)
.map_err(ConnectError::Connect)?;
connecting.await.map_err(ConnectError::Connection)
}
fn resolve(host: &str, port: u16) -> Result<SocketAddr, ConnectError> {
(host, port)
.to_socket_addrs()
.map_err(ConnectError::Resolve)?
.next()
.ok_or(ConnectError::NoAddress)
}