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
30#[derive(Clone, Copy)]
40pub enum Trust {
41 Pinned([u8; 32]),
46 WebPki,
50 Insecure,
53}
54
55#[derive(Debug)]
56pub enum ConnectError {
57 Resolve(std::io::Error),
58 NoAddress,
59 Endpoint(std::io::Error),
60 Config(rustls::Error),
61 Connect(quinn::ConnectError),
62 Connection(quinn::ConnectionError),
63}
64
65impl std::fmt::Display for ConnectError {
66 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
67 match self {
68 ConnectError::Resolve(e) => write!(f, "resolving station address: {e}"),
69 ConnectError::NoAddress => write!(f, "hostname resolved to no addresses"),
70 ConnectError::Endpoint(e) => write!(f, "creating QUIC endpoint: {e}"),
71 ConnectError::Config(e) => write!(f, "building TLS config: {e}"),
72 ConnectError::Connect(e) => write!(f, "starting QUIC connect: {e}"),
73 ConnectError::Connection(e) => write!(f, "QUIC connection failed: {e}"),
74 }
75 }
76}
77
78impl std::error::Error for ConnectError {}
79
80fn client_config(trust: Trust) -> Result<ClientConfig, rustls::Error> {
81 let mut crypto = match trust {
82 Trust::Pinned(pubkey) => rustls::ClientConfig::builder()
83 .dangerous()
84 .with_custom_certificate_verifier(Arc::new(PubkeyPinVerifier::new(pubkey)))
85 .with_no_client_auth(),
86 Trust::WebPki => {
87 let mut roots = rustls::RootCertStore::empty();
88 roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
89 rustls::ClientConfig::builder()
90 .with_root_certificates(roots)
91 .with_no_client_auth()
92 }
93 Trust::Insecure => rustls::ClientConfig::builder()
94 .dangerous()
95 .with_custom_certificate_verifier(Arc::new(SkipServerVerification::new()))
96 .with_no_client_auth(),
97 };
98 crypto.alpn_protocols = vec![ALPN.to_vec()];
99
100 let mut transport = TransportConfig::default();
101 transport.max_idle_timeout(Some(
102 IdleTimeout::try_from(DEFAULT_IDLE_TIMEOUT).expect("valid idle timeout"),
103 ));
104 transport.keep_alive_interval(Some(DEFAULT_KEEP_ALIVE_INTERVAL));
105 apply_flow_control_defaults(&mut transport);
106
107 let quic_crypto = quinn::crypto::rustls::QuicClientConfig::try_from(crypto)
108 .map_err(|e| rustls::Error::General(e.to_string()))?;
109 let mut config = ClientConfig::new(Arc::new(quic_crypto));
110 config.transport_config(Arc::new(transport));
111 Ok(config)
112}
113
114fn apply_flow_control_defaults(transport: &mut TransportConfig) {
120 transport.stream_receive_window((16u32 * 1024 * 1024).into());
121 transport.receive_window((64u32 * 1024 * 1024).into());
122 transport.send_window(64u64 * 1024 * 1024);
123}
124
125pub async fn connect(
131 host: &str,
132 port: u16,
133 trust: Trust,
134) -> Result<quinn::Connection, ConnectError> {
135 let addr = resolve(host, port)?;
136 let bind_addr: SocketAddr = if addr.is_ipv6() {
137 "[::]:0".parse().expect("valid unspecified v6 addr")
138 } else {
139 "0.0.0.0:0".parse().expect("valid unspecified v4 addr")
140 };
141
142 let mut endpoint = Endpoint::client(bind_addr).map_err(ConnectError::Endpoint)?;
143 let config = client_config(trust).map_err(ConnectError::Config)?;
144 endpoint.set_default_client_config(config);
145
146 let connecting = endpoint
147 .connect(addr, host)
148 .map_err(ConnectError::Connect)?;
149 connecting.await.map_err(ConnectError::Connection)
150}
151
152fn resolve(host: &str, port: u16) -> Result<SocketAddr, ConnectError> {
153 (host, port)
154 .to_socket_addrs()
155 .map_err(ConnectError::Resolve)?
156 .next()
157 .ok_or(ConnectError::NoAddress)
158}