1use std::net::{SocketAddr, ToSocketAddrs};
27use std::sync::Arc;
28use std::time::Duration;
29
30use macula_pqc::KeyPossessionVerifier;
31use quinn::crypto::rustls::QuicClientConfig;
32use quinn::{ClientConfig, Endpoint, IdleTimeout, TransportConfig};
33
34use crate::profile::Profile;
35
36pub const ALPN: &[u8] = b"macula";
38
39pub const IDLE_TIMEOUT: Duration = Duration::from_secs(300);
43pub const KEEP_ALIVE_INTERVAL: Duration = Duration::from_secs(15);
44
45#[derive(Debug, Clone, PartialEq, Eq)]
48pub struct Target {
49 pub host: String,
50 pub port: u16,
51 pub profile: Profile,
52 pub expected_node_id: [u8; 32],
53}
54
55pub struct Dialed {
59 pub connection: quinn::Connection,
60 pub endpoint: Endpoint,
61 pub leaf: Vec<u8>,
62 pub target: Target,
63}
64
65#[derive(Debug)]
67pub enum DialError {
68 NoExpectedNodeId,
70 Resolve(std::io::Error),
72 Endpoint(std::io::Error),
74 Config(String),
76 Connect(quinn::ConnectError),
78 Connection(quinn::ConnectionError),
82 NoLeaf,
84}
85
86impl std::fmt::Display for DialError {
87 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
88 match self {
89 DialError::NoExpectedNodeId => f.write_str("the dial target names no expected node_id"),
90 DialError::Resolve(e) => write!(f, "resolving the station's address: {e}"),
91 DialError::Endpoint(e) => write!(f, "creating the QUIC endpoint: {e}"),
92 DialError::Config(e) => write!(f, "building the TLS configuration: {e}"),
93 DialError::Connect(e) => write!(f, "starting the QUIC connection: {e}"),
94 DialError::Connection(e) => write!(f, "the QUIC connection failed: {e}"),
95 DialError::NoLeaf => f.write_str("the station presented no certificate"),
96 }
97 }
98}
99
100impl std::error::Error for DialError {}
101
102pub async fn dial_target(target: &Target) -> Result<Dialed, DialError> {
105 if target.expected_node_id == [0u8; 32] {
106 return Err(DialError::NoExpectedNodeId);
107 }
108 let addr = (target.host.as_str(), target.port)
109 .to_socket_addrs()
110 .map_err(DialError::Resolve)?
111 .next()
112 .ok_or_else(|| DialError::Resolve(std::io::Error::other("no address")))?;
113 let bind: SocketAddr = if addr.is_ipv6() {
114 (std::net::Ipv6Addr::UNSPECIFIED, 0).into()
115 } else {
116 (std::net::Ipv4Addr::UNSPECIFIED, 0).into()
117 };
118 let mut endpoint = Endpoint::client(bind).map_err(DialError::Endpoint)?;
119 endpoint.set_default_client_config(client_config()?);
120 let connection = endpoint
121 .connect(addr, &target.host)
122 .map_err(DialError::Connect)?
123 .await
124 .map_err(DialError::Connection)?;
125 let leaf = connection
126 .peer_identity()
127 .and_then(|identity| {
128 identity
129 .downcast::<Vec<rustls::pki_types::CertificateDer<'static>>>()
130 .ok()
131 })
132 .and_then(|chain| chain.first().map(|leaf| leaf.to_vec()))
133 .ok_or(DialError::NoLeaf)?;
134 Ok(Dialed {
135 connection,
136 endpoint,
137 leaf,
138 target: target.clone(),
139 })
140}
141
142fn client_config() -> Result<ClientConfig, DialError> {
144 let quic = QuicClientConfig::with_initial(
145 Arc::new(tls_client_config()?),
146 macula_pqc::quic_initial_suite(),
147 )
148 .map_err(|e| DialError::Config(e.to_string()))?;
149 let mut transport = TransportConfig::default();
150 transport.max_idle_timeout(Some(
151 IdleTimeout::try_from(IDLE_TIMEOUT).map_err(|e| DialError::Config(e.to_string()))?,
152 ));
153 transport.keep_alive_interval(Some(KEEP_ALIVE_INTERVAL));
154 transport.stream_receive_window((16u32 * 1024 * 1024).into());
155 transport.receive_window((64u32 * 1024 * 1024).into());
156 transport.send_window(64 * 1024 * 1024);
157 let mut config = ClientConfig::new(Arc::new(quic));
158 config.transport_config(Arc::new(transport));
159 Ok(config)
160}
161
162const SECP384R1MLKEM1024: rustls::NamedGroup = rustls::NamedGroup::Unknown(0x11ED);
164
165pub(crate) fn tls_client_config() -> Result<rustls::ClientConfig, DialError> {
169 let verifier = KeyPossessionVerifier::new();
170 let pqc = macula_pqc::client_builder();
171 let kx_groups: Vec<_> = pqc
172 .crypto_provider()
173 .kx_groups
174 .iter()
175 .copied()
176 .filter(|group| group.name() == SECP384R1MLKEM1024)
177 .collect();
178 if kx_groups.len() != 1 {
179 return Err(DialError::Config(format!(
180 "macula-pqc's provider carries {} SecP384r1MLKEM1024 groups, not 1",
181 kx_groups.len()
182 )));
183 }
184 let provider = rustls::crypto::CryptoProvider {
185 kx_groups,
186 ..(**pqc.crypto_provider()).clone()
187 };
188 let mut config = rustls::ClientConfig::builder_with_provider(Arc::new(provider))
189 .with_protocol_versions(&[&rustls::version::TLS13])
190 .map_err(|e| DialError::Config(e.to_string()))?
191 .dangerous()
192 .with_custom_certificate_verifier(Arc::new(verifier))
193 .with_no_client_auth();
194 config.alpn_protocols = vec![ALPN.to_vec()];
195 config.resumption = rustls::client::Resumption::disabled();
196 config.enable_early_data = false;
197 Ok(config)
198}
199
200#[cfg(test)]
201mod tests {
202 use std::sync::Arc;
207
208 use quinn::crypto::rustls::QuicServerConfig;
209 use rustls::crypto::aws_lc_rs::kx_group as aws;
210 use rustls::crypto::CryptoProvider;
211 use rustls::pki_types::{CertificateDer, PrivateKeyDer};
212 use rustls::{NamedGroup, ServerConfig};
213
214 use super::{dial_target, tls_client_config, DialError, Target, ALPN};
215 use crate::profile::Profile;
216
217 const SECP384R1MLKEM1024: NamedGroup = NamedGroup::Unknown(0x11ED);
219
220 fn mldsa_certificate() -> (CertificateDer<'static>, PrivateKeyDer<'static>) {
221 let (certificate, key) =
222 macula_pqc::self_signed_certificate(&[7u8; 32], vec!["localhost".to_string()])
223 .expect("a certificate");
224 (certificate, key.into())
225 }
226
227 fn classical_certificate() -> (CertificateDer<'static>, PrivateKeyDer<'static>) {
228 let key_pair = rcgen::KeyPair::generate().expect("a key pair");
229 let certificate = rcgen::CertificateParams::new(vec!["localhost".to_string()])
230 .expect("certificate params")
231 .self_signed(&key_pair)
232 .expect("a certificate");
233 (
234 certificate.der().clone(),
235 PrivateKeyDer::Pkcs8(key_pair.serialize_der().into()),
236 )
237 }
238
239 fn station(mut config: ServerConfig, alpn: &[u8]) -> (quinn::Endpoint, u16) {
242 config.alpn_protocols = vec![alpn.to_vec()];
243 let quic =
244 QuicServerConfig::with_initial(Arc::new(config), macula_pqc::quic_initial_suite())
245 .expect("a QUIC server config");
246 let endpoint = quinn::Endpoint::server(
247 quinn::ServerConfig::with_crypto(Arc::new(quic)),
248 ([127, 0, 0, 1], 0).into(),
249 )
250 .expect("an endpoint");
251 let port = endpoint.local_addr().expect("an address").port();
252 let accepting = endpoint.clone();
253 tokio::spawn(hold_connections(accepting));
254 (endpoint, port)
255 }
256
257 async fn hold_connections(accepting: quinn::Endpoint) {
260 while let Some(incoming) = accepting.accept().await {
261 tokio::spawn(hold_until_closed(incoming));
262 }
263 }
264
265 async fn hold_until_closed(incoming: quinn::Incoming) {
267 if let Ok(connection) = incoming.await {
268 connection.closed().await;
269 }
270 }
271
272 fn macula_station() -> (quinn::Endpoint, u16, CertificateDer<'static>) {
273 let (certificate, key) = mldsa_certificate();
274 let config = macula_pqc::server_builder()
275 .with_no_client_auth()
276 .with_single_cert(vec![certificate.clone()], key)
277 .expect("a station configuration");
278 let (endpoint, port) = station(config, ALPN);
279 (endpoint, port, certificate)
280 }
281
282 fn target(port: u16) -> Target {
283 Target {
284 host: "127.0.0.1".to_string(),
285 port,
286 profile: Profile::PqHybrid,
287 expected_node_id: [1u8; 32],
288 }
289 }
290
291 #[test]
292 fn a_dial_offers_secp384r1mlkem1024_alone() {
293 let config = tls_client_config().expect("a configuration");
294 let offered: Vec<NamedGroup> = config
295 .crypto_provider()
296 .kx_groups
297 .iter()
298 .map(|g| g.name())
299 .collect();
300 assert_eq!(offered, vec![SECP384R1MLKEM1024]);
301 }
302
303 fn station_on_group(group: NamedGroup) -> (quinn::Endpoint, u16) {
307 let (certificate, key) = mldsa_certificate();
308 let pqc = macula_pqc::server_builder();
309 let provider = CryptoProvider {
310 kx_groups: pqc
311 .crypto_provider()
312 .kx_groups
313 .iter()
314 .copied()
315 .filter(|g| g.name() == group)
316 .collect(),
317 ..(**pqc.crypto_provider()).clone()
318 };
319 assert_eq!(provider.kx_groups.len(), 1, "macula-pqc offers {group:?}");
320 let config = ServerConfig::builder_with_provider(Arc::new(provider))
321 .with_protocol_versions(&[&rustls::version::TLS13])
322 .expect("TLS 1.3")
323 .with_no_client_auth()
324 .with_single_cert(vec![certificate], key)
325 .expect("a station configuration");
326 station(config, ALPN)
327 }
328
329 #[tokio::test]
330 async fn a_station_that_offers_only_secp256r1mlkem768_is_refused() {
331 let (_station, port) = station_on_group(NamedGroup::secp256r1MLKEM768);
332 assert!(matches!(
333 dial_target(&target(port)).await,
334 Err(DialError::Connection(_))
335 ));
336 }
337
338 #[tokio::test]
339 async fn a_station_that_offers_only_secp384r1mlkem1024_is_reached() {
340 let (_station, port) = station_on_group(SECP384R1MLKEM1024);
341 let dialed = dial_target(&target(port))
342 .await
343 .expect("the station is reached");
344 dialed.connection.close(0u32.into(), b"done");
345 }
346
347 #[tokio::test]
348 async fn a_macula_12_station_is_reached_and_its_leaf_handed_back() {
349 let (_station, port, certificate) = macula_station();
350 let dialed = dial_target(&target(port))
351 .await
352 .expect("the station is reached");
353 assert_eq!(dialed.leaf, certificate.to_vec());
354 dialed.connection.close(0u32.into(), b"done");
355 }
356
357 #[tokio::test]
358 async fn a_target_without_an_expected_node_id_is_refused_before_dialing() {
359 let mut unpinned = target(9);
360 unpinned.expected_node_id = [0u8; 32];
361 assert!(matches!(
362 dial_target(&unpinned).await,
363 Err(DialError::NoExpectedNodeId)
364 ));
365 }
366
367 #[tokio::test]
368 async fn a_station_with_a_classical_certificate_is_refused() {
369 let (certificate, key) = classical_certificate();
370 let provider = CryptoProvider {
371 kx_groups: vec![aws::SECP256R1MLKEM768],
372 ..rustls::crypto::aws_lc_rs::default_provider()
373 };
374 let config = ServerConfig::builder_with_provider(Arc::new(provider))
375 .with_protocol_versions(&[&rustls::version::TLS13])
376 .expect("TLS 1.3")
377 .with_no_client_auth()
378 .with_single_cert(vec![certificate], key)
379 .expect("a station configuration");
380 let (_station, port) = station(config, ALPN);
381 assert!(matches!(
382 dial_target(&target(port)).await,
383 Err(DialError::Connection(_))
384 ));
385 }
386
387 #[tokio::test]
388 async fn a_classical_only_station_is_refused() {
389 let (certificate, key) = classical_certificate();
390 let provider = CryptoProvider {
391 kx_groups: vec![aws::X25519, aws::SECP256R1, aws::SECP384R1],
392 ..rustls::crypto::aws_lc_rs::default_provider()
393 };
394 let config = ServerConfig::builder_with_provider(Arc::new(provider))
395 .with_protocol_versions(&[&rustls::version::TLS13])
396 .expect("TLS 1.3")
397 .with_no_client_auth()
398 .with_single_cert(vec![certificate], key)
399 .expect("a station configuration");
400 let (_station, port) = station(config, ALPN);
401 assert!(matches!(
402 dial_target(&target(port)).await,
403 Err(DialError::Connection(_))
404 ));
405 }
406
407 #[tokio::test]
408 async fn a_station_that_does_not_speak_macula_is_refused() {
409 let (certificate, key) = mldsa_certificate();
410 let config = macula_pqc::server_builder()
411 .with_no_client_auth()
412 .with_single_cert(vec![certificate], key)
413 .expect("a station configuration");
414 let (_station, port) = station(config, b"h3");
415 assert!(matches!(
416 dial_target(&target(port)).await,
417 Err(DialError::Connection(_))
418 ));
419 }
420}