1use std::net::{SocketAddr, ToSocketAddrs};
20use std::sync::Arc;
21use std::time::Duration;
22
23use macula_pqc::KeyPossessionVerifier;
24use quinn::crypto::rustls::QuicClientConfig;
25use quinn::{ClientConfig, Endpoint, IdleTimeout, TransportConfig};
26
27use crate::profile::Profile;
28
29pub const ALPN: &[u8] = b"macula";
31
32pub const IDLE_TIMEOUT: Duration = Duration::from_secs(300);
36pub const KEEP_ALIVE_INTERVAL: Duration = Duration::from_secs(15);
37
38#[derive(Debug, Clone, PartialEq, Eq)]
41pub struct Target {
42 pub host: String,
43 pub port: u16,
44 pub profile: Profile,
45 pub expected_node_id: [u8; 32],
46}
47
48pub struct Dialed {
52 pub connection: quinn::Connection,
53 pub endpoint: Endpoint,
54 pub leaf: Vec<u8>,
55 pub target: Target,
56}
57
58#[derive(Debug)]
60pub enum DialError {
61 NoExpectedNodeId,
63 Resolve(std::io::Error),
65 Endpoint(std::io::Error),
67 Config(String),
69 Connect(quinn::ConnectError),
71 Connection(quinn::ConnectionError),
74 NoLeaf,
76}
77
78impl std::fmt::Display for DialError {
79 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
80 match self {
81 DialError::NoExpectedNodeId => f.write_str("the dial target names no expected node_id"),
82 DialError::Resolve(e) => write!(f, "resolving the station's address: {e}"),
83 DialError::Endpoint(e) => write!(f, "creating the QUIC endpoint: {e}"),
84 DialError::Config(e) => write!(f, "building the TLS configuration: {e}"),
85 DialError::Connect(e) => write!(f, "starting the QUIC connection: {e}"),
86 DialError::Connection(e) => write!(f, "the QUIC connection failed: {e}"),
87 DialError::NoLeaf => f.write_str("the station presented no certificate"),
88 }
89 }
90}
91
92impl std::error::Error for DialError {}
93
94pub async fn dial_target(target: &Target) -> Result<Dialed, DialError> {
97 if target.expected_node_id == [0u8; 32] {
98 return Err(DialError::NoExpectedNodeId);
99 }
100 let addr = (target.host.as_str(), target.port)
101 .to_socket_addrs()
102 .map_err(DialError::Resolve)?
103 .next()
104 .ok_or_else(|| DialError::Resolve(std::io::Error::other("no address")))?;
105 let bind: SocketAddr = if addr.is_ipv6() {
106 (std::net::Ipv6Addr::UNSPECIFIED, 0).into()
107 } else {
108 (std::net::Ipv4Addr::UNSPECIFIED, 0).into()
109 };
110 let mut endpoint = Endpoint::client(bind).map_err(DialError::Endpoint)?;
111 endpoint.set_default_client_config(client_config()?);
112 let connection = endpoint
113 .connect(addr, &target.host)
114 .map_err(DialError::Connect)?
115 .await
116 .map_err(DialError::Connection)?;
117 let leaf = connection
118 .peer_identity()
119 .and_then(|identity| {
120 identity
121 .downcast::<Vec<rustls::pki_types::CertificateDer<'static>>>()
122 .ok()
123 })
124 .and_then(|chain| chain.first().map(|leaf| leaf.to_vec()))
125 .ok_or(DialError::NoLeaf)?;
126 Ok(Dialed {
127 connection,
128 endpoint,
129 leaf,
130 target: target.clone(),
131 })
132}
133
134fn client_config() -> Result<ClientConfig, DialError> {
136 let quic = QuicClientConfig::with_initial(
137 Arc::new(tls_client_config()?),
138 macula_pqc::quic_initial_suite(),
139 )
140 .map_err(|e| DialError::Config(e.to_string()))?;
141 let mut transport = TransportConfig::default();
142 transport.max_idle_timeout(Some(
143 IdleTimeout::try_from(IDLE_TIMEOUT).map_err(|e| DialError::Config(e.to_string()))?,
144 ));
145 transport.keep_alive_interval(Some(KEEP_ALIVE_INTERVAL));
146 transport.stream_receive_window((16u32 * 1024 * 1024).into());
147 transport.receive_window((64u32 * 1024 * 1024).into());
148 transport.send_window(64 * 1024 * 1024);
149 let mut config = ClientConfig::new(Arc::new(quic));
150 config.transport_config(Arc::new(transport));
151 Ok(config)
152}
153
154pub(crate) fn tls_client_config() -> Result<rustls::ClientConfig, DialError> {
157 let verifier = KeyPossessionVerifier::new();
158 let mut config = macula_pqc::client_builder()
159 .dangerous()
160 .with_custom_certificate_verifier(Arc::new(verifier))
161 .with_no_client_auth();
162 config.alpn_protocols = vec![ALPN.to_vec()];
163 config.resumption = rustls::client::Resumption::disabled();
164 config.enable_early_data = false;
165 Ok(config)
166}
167
168#[cfg(test)]
169mod tests {
170 use std::sync::Arc;
175
176 use quinn::crypto::rustls::QuicServerConfig;
177 use rustls::crypto::aws_lc_rs::kx_group as aws;
178 use rustls::crypto::CryptoProvider;
179 use rustls::pki_types::{CertificateDer, PrivateKeyDer};
180 use rustls::{NamedGroup, ServerConfig};
181
182 use super::{dial_target, tls_client_config, DialError, Target, ALPN};
183 use crate::profile::Profile;
184
185 const SECP384R1MLKEM1024: NamedGroup = NamedGroup::Unknown(0x11ED);
187
188 fn mldsa_certificate() -> (CertificateDer<'static>, PrivateKeyDer<'static>) {
189 let (certificate, key) =
190 macula_pqc::self_signed_certificate(&[7u8; 32], vec!["localhost".to_string()])
191 .expect("a certificate");
192 (certificate, key.into())
193 }
194
195 fn classical_certificate() -> (CertificateDer<'static>, PrivateKeyDer<'static>) {
196 let key_pair = rcgen::KeyPair::generate().expect("a key pair");
197 let certificate = rcgen::CertificateParams::new(vec!["localhost".to_string()])
198 .expect("certificate params")
199 .self_signed(&key_pair)
200 .expect("a certificate");
201 (
202 certificate.der().clone(),
203 PrivateKeyDer::Pkcs8(key_pair.serialize_der().into()),
204 )
205 }
206
207 fn station(mut config: ServerConfig, alpn: &[u8]) -> (quinn::Endpoint, u16) {
210 config.alpn_protocols = vec![alpn.to_vec()];
211 let quic =
212 QuicServerConfig::with_initial(Arc::new(config), macula_pqc::quic_initial_suite())
213 .expect("a QUIC server config");
214 let endpoint = quinn::Endpoint::server(
215 quinn::ServerConfig::with_crypto(Arc::new(quic)),
216 ([127, 0, 0, 1], 0).into(),
217 )
218 .expect("an endpoint");
219 let port = endpoint.local_addr().expect("an address").port();
220 let accepting = endpoint.clone();
221 tokio::spawn(async move {
222 while let Some(incoming) = accepting.accept().await {
223 tokio::spawn(async move {
224 if let Ok(connection) = incoming.await {
225 connection.closed().await;
226 }
227 });
228 }
229 });
230 (endpoint, port)
231 }
232
233 fn macula_station() -> (quinn::Endpoint, u16, CertificateDer<'static>) {
234 let (certificate, key) = mldsa_certificate();
235 let config = macula_pqc::server_builder()
236 .with_no_client_auth()
237 .with_single_cert(vec![certificate.clone()], key)
238 .expect("a station configuration");
239 let (endpoint, port) = station(config, ALPN);
240 (endpoint, port, certificate)
241 }
242
243 fn target(port: u16) -> Target {
244 Target {
245 host: "127.0.0.1".to_string(),
246 port,
247 profile: Profile::PqHybrid,
248 expected_node_id: [1u8; 32],
249 }
250 }
251
252 #[test]
253 fn a_dial_offers_exactly_macula_pqcs_groups() {
254 let config = tls_client_config().expect("a configuration");
255 let offered: Vec<NamedGroup> = config
256 .crypto_provider()
257 .kx_groups
258 .iter()
259 .map(|g| g.name())
260 .collect();
261 assert_eq!(
262 offered,
263 vec![SECP384R1MLKEM1024, NamedGroup::secp256r1MLKEM768]
264 );
265 }
266
267 #[tokio::test]
268 async fn a_macula_12_station_is_reached_and_its_leaf_handed_back() {
269 let (_station, port, certificate) = macula_station();
270 let dialed = dial_target(&target(port))
271 .await
272 .expect("the station is reached");
273 assert_eq!(dialed.leaf, certificate.to_vec());
274 dialed.connection.close(0u32.into(), b"done");
275 }
276
277 #[tokio::test]
278 async fn a_target_without_an_expected_node_id_is_refused_before_dialing() {
279 let mut unpinned = target(9);
280 unpinned.expected_node_id = [0u8; 32];
281 assert!(matches!(
282 dial_target(&unpinned).await,
283 Err(DialError::NoExpectedNodeId)
284 ));
285 }
286
287 #[tokio::test]
288 async fn a_station_with_a_classical_certificate_is_refused() {
289 let (certificate, key) = classical_certificate();
290 let provider = CryptoProvider {
291 kx_groups: vec![aws::SECP256R1MLKEM768],
292 ..rustls::crypto::aws_lc_rs::default_provider()
293 };
294 let config = ServerConfig::builder_with_provider(Arc::new(provider))
295 .with_protocol_versions(&[&rustls::version::TLS13])
296 .expect("TLS 1.3")
297 .with_no_client_auth()
298 .with_single_cert(vec![certificate], key)
299 .expect("a station configuration");
300 let (_station, port) = station(config, ALPN);
301 assert!(matches!(
302 dial_target(&target(port)).await,
303 Err(DialError::Connection(_))
304 ));
305 }
306
307 #[tokio::test]
308 async fn a_classical_only_station_is_refused() {
309 let (certificate, key) = classical_certificate();
310 let provider = CryptoProvider {
311 kx_groups: vec![aws::X25519, aws::SECP256R1, aws::SECP384R1],
312 ..rustls::crypto::aws_lc_rs::default_provider()
313 };
314 let config = ServerConfig::builder_with_provider(Arc::new(provider))
315 .with_protocol_versions(&[&rustls::version::TLS13])
316 .expect("TLS 1.3")
317 .with_no_client_auth()
318 .with_single_cert(vec![certificate], key)
319 .expect("a station configuration");
320 let (_station, port) = station(config, ALPN);
321 assert!(matches!(
322 dial_target(&target(port)).await,
323 Err(DialError::Connection(_))
324 ));
325 }
326
327 #[tokio::test]
328 async fn a_station_that_does_not_speak_macula_is_refused() {
329 let (certificate, key) = mldsa_certificate();
330 let config = macula_pqc::server_builder()
331 .with_no_client_auth()
332 .with_single_cert(vec![certificate], key)
333 .expect("a station configuration");
334 let (_station, port) = station(config, b"h3");
335 assert!(matches!(
336 dial_target(&target(port)).await,
337 Err(DialError::Connection(_))
338 ));
339 }
340}