1use std::net::{SocketAddr, ToSocketAddrs};
12use std::sync::Arc;
13use std::time::Duration;
14
15use quinn::{ClientConfig, Endpoint, IdleTimeout, TransportConfig};
16
17use crate::cert::{PubkeyPinVerifier, SkipServerVerification};
18
19pub const ALPN: &[u8] = b"macula";
21
22pub const DEFAULT_IDLE_TIMEOUT: Duration = Duration::from_secs(300);
28pub const DEFAULT_KEEP_ALIVE_INTERVAL: Duration = Duration::from_secs(15);
29
30pub enum Trust {
34 Pinned([u8; 32]),
39 WebPki,
43 Insecure,
46}
47
48#[derive(Debug)]
49pub enum ConnectError {
50 Resolve(std::io::Error),
51 NoAddress,
52 Endpoint(std::io::Error),
53 Config(rustls::Error),
54 Connect(quinn::ConnectError),
55 Connection(quinn::ConnectionError),
56}
57
58impl std::fmt::Display for ConnectError {
59 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
60 match self {
61 ConnectError::Resolve(e) => write!(f, "resolving station address: {e}"),
62 ConnectError::NoAddress => write!(f, "hostname resolved to no addresses"),
63 ConnectError::Endpoint(e) => write!(f, "creating QUIC endpoint: {e}"),
64 ConnectError::Config(e) => write!(f, "building TLS config: {e}"),
65 ConnectError::Connect(e) => write!(f, "starting QUIC connect: {e}"),
66 ConnectError::Connection(e) => write!(f, "QUIC connection failed: {e}"),
67 }
68 }
69}
70
71impl std::error::Error for ConnectError {}
72
73fn client_config(trust: Trust) -> Result<ClientConfig, rustls::Error> {
74 let mut crypto = match trust {
75 Trust::Pinned(pubkey) => rustls::ClientConfig::builder()
76 .dangerous()
77 .with_custom_certificate_verifier(Arc::new(PubkeyPinVerifier::new(pubkey)))
78 .with_no_client_auth(),
79 Trust::WebPki => {
80 let mut roots = rustls::RootCertStore::empty();
81 roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
82 rustls::ClientConfig::builder()
83 .with_root_certificates(roots)
84 .with_no_client_auth()
85 }
86 Trust::Insecure => rustls::ClientConfig::builder()
87 .dangerous()
88 .with_custom_certificate_verifier(Arc::new(SkipServerVerification::new()))
89 .with_no_client_auth(),
90 };
91 crypto.alpn_protocols = vec![ALPN.to_vec()];
92
93 let mut transport = TransportConfig::default();
94 transport.max_idle_timeout(Some(
95 IdleTimeout::try_from(DEFAULT_IDLE_TIMEOUT).expect("valid idle timeout"),
96 ));
97 transport.keep_alive_interval(Some(DEFAULT_KEEP_ALIVE_INTERVAL));
98 apply_flow_control_defaults(&mut transport);
99
100 let quic_crypto = quinn::crypto::rustls::QuicClientConfig::try_from(crypto)
101 .map_err(|e| rustls::Error::General(e.to_string()))?;
102 let mut config = ClientConfig::new(Arc::new(quic_crypto));
103 config.transport_config(Arc::new(transport));
104 Ok(config)
105}
106
107fn apply_flow_control_defaults(transport: &mut TransportConfig) {
113 transport.stream_receive_window((16u32 * 1024 * 1024).into());
114 transport.receive_window((64u32 * 1024 * 1024).into());
115 transport.send_window(64u64 * 1024 * 1024);
116}
117
118pub async fn connect(
124 host: &str,
125 port: u16,
126 trust: Trust,
127) -> Result<quinn::Connection, ConnectError> {
128 let addr = resolve(host, port)?;
129 let bind_addr: SocketAddr = if addr.is_ipv6() {
130 "[::]:0".parse().expect("valid unspecified v6 addr")
131 } else {
132 "0.0.0.0:0".parse().expect("valid unspecified v4 addr")
133 };
134
135 let mut endpoint = Endpoint::client(bind_addr).map_err(ConnectError::Endpoint)?;
136 let config = client_config(trust).map_err(ConnectError::Config)?;
137 endpoint.set_default_client_config(config);
138
139 let connecting = endpoint
140 .connect(addr, host)
141 .map_err(ConnectError::Connect)?;
142 connecting.await.map_err(ConnectError::Connection)
143}
144
145fn resolve(host: &str, port: u16) -> Result<SocketAddr, ConnectError> {
146 (host, port)
147 .to_socket_addrs()
148 .map_err(ConnectError::Resolve)?
149 .next()
150 .ok_or(ConnectError::NoAddress)
151}