Skip to main content

rustls_post_quantum/
lib.rs

1//! This crate provide a [`CryptoProvider`] built on the default aws-lc-rs default provider.
2//!
3//! Features:
4//!
5//! - `aws-lc-rs-unstable`: adds support for three variants of the experimental ML-DSA signature
6//!   algorithm.
7//!
8//! Before rustls 0.23.22, this crate additionally provided support for the ML-KEM key exchange
9//! (both "pure" and hybrid variants), but these have been moved to the rustls crate itself.
10//! In rustls 0.23.22 and later, you can use rustls' `prefer-post-quantum` feature to determine
11//! whether the ML-KEM key exchange is preferred over non-post-quantum key exchanges.
12
13use core::fmt::{self, Debug, Formatter};
14use std::sync::Arc;
15
16use aws_lc_rs::signature::KeyPair;
17use aws_lc_rs::unstable::signature::{
18    ML_DSA_44_SIGNING, ML_DSA_65_SIGNING, ML_DSA_87_SIGNING, PqdsaKeyPair, PqdsaSigningAlgorithm,
19};
20use rustls::Error;
21use rustls::crypto::{
22    CryptoProvider, KeyProvider, SignatureScheme, Signer, SigningKey, WebPkiSupportedAlgorithms,
23    public_key_to_spki,
24};
25use rustls::pki_types::{
26    AlgorithmIdentifier, FipsStatus, PrivateKeyDer, SignatureVerificationAlgorithm,
27    SubjectPublicKeyInfoDer, alg_id,
28};
29use rustls_aws_lc_rs::AwsLcRsVerificationAlgorithm;
30
31/// The default `CryptoProvider` backed by aws-lc-rs.
32pub const DEFAULT_PROVIDER: CryptoProvider = CryptoProvider {
33    signature_verification_algorithms: SUPPORTED_SIG_ALGS,
34    key_provider: &PqAwsLcRs,
35    ..rustls_aws_lc_rs::DEFAULT_PROVIDER
36};
37
38#[derive(Debug)]
39pub struct PqAwsLcRs;
40
41impl KeyProvider for PqAwsLcRs {
42    fn load_private_key(
43        &self,
44        key_der: PrivateKeyDer<'static>,
45    ) -> Result<Box<dyn SigningKey>, Error> {
46        // TODO: support `PqdsaKeyPair::from_raw_private_key()`?
47        if let PrivateKeyDer::Pkcs8(pkcs8) = &key_der {
48            for kind in PqdsaKeyKind::iter() {
49                match PqdsaKeyPair::from_pkcs8(kind.to_alg(), pkcs8.secret_pkcs8_der()) {
50                    Ok(key_pair) => {
51                        return Ok(Box::new(PqdsaSigningKey {
52                            kind,
53                            inner: Arc::new(key_pair),
54                        }));
55                    }
56                    Err(_) => continue,
57                }
58            }
59        }
60
61        match rustls_aws_lc_rs::DEFAULT_KEY_PROVIDER.load_private_key(key_der) {
62            Ok(key) => Ok(key),
63            Err(_) => Err(Error::General(
64                "failed to parse private key as ML-DSA, RSA, ECDSA, or EdDSA".into(),
65            )),
66        }
67    }
68
69    fn fips(&self) -> FipsStatus {
70        FipsStatus::Unvalidated
71    }
72}
73
74struct PqdsaSigningKey {
75    kind: PqdsaKeyKind,
76    inner: Arc<PqdsaKeyPair>,
77}
78
79impl SigningKey for PqdsaSigningKey {
80    fn choose_scheme(&self, offered: &[SignatureScheme]) -> Option<Box<dyn Signer>> {
81        if !offered.contains(&self.kind.scheme()) {
82            return None;
83        }
84
85        Some(Box::new(PqdsaSigner {
86            key: self.inner.clone(),
87            kind: self.kind,
88        }))
89    }
90
91    fn public_key(&self) -> Option<SubjectPublicKeyInfoDer<'_>> {
92        Some(public_key_to_spki(
93            &self.kind.alg_id(),
94            self.inner.public_key(),
95        ))
96    }
97}
98
99impl Debug for PqdsaSigningKey {
100    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
101        f.debug_struct("PqdsaSigningKey")
102            .field("scheme", &self.kind.scheme())
103            .finish_non_exhaustive()
104    }
105}
106
107struct PqdsaSigner {
108    key: Arc<PqdsaKeyPair>,
109    kind: PqdsaKeyKind,
110}
111
112impl Signer for PqdsaSigner {
113    fn sign(self: Box<Self>, message: &[u8]) -> Result<Vec<u8>, Error> {
114        let expected_sig_len = self.key.algorithm().signature_len();
115        let mut sig = vec![0; expected_sig_len];
116        let actual_sig_len = self
117            .key
118            .sign(message, &mut sig)
119            .map_err(|_| Error::General("signing failed".into()))?;
120
121        if actual_sig_len != expected_sig_len {
122            return Err(Error::General("unexpected signature length".into()));
123        }
124
125        Ok(sig)
126    }
127
128    fn scheme(&self) -> SignatureScheme {
129        self.kind.scheme()
130    }
131}
132
133impl Debug for PqdsaSigner {
134    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
135        f.debug_struct("PqdsaSigner")
136            .field("scheme", &self.kind.scheme())
137            .finish_non_exhaustive()
138    }
139}
140
141#[derive(Clone, Copy)]
142enum PqdsaKeyKind {
143    MlDsa44,
144    MlDsa65,
145    MlDsa87,
146}
147
148impl PqdsaKeyKind {
149    fn iter() -> impl Iterator<Item = Self> {
150        [Self::MlDsa44, Self::MlDsa65, Self::MlDsa87].into_iter()
151    }
152
153    fn to_alg(self) -> &'static PqdsaSigningAlgorithm {
154        match self {
155            Self::MlDsa44 => &ML_DSA_44_SIGNING,
156            Self::MlDsa65 => &ML_DSA_65_SIGNING,
157            Self::MlDsa87 => &ML_DSA_87_SIGNING,
158        }
159    }
160
161    fn scheme(&self) -> SignatureScheme {
162        match self {
163            Self::MlDsa44 => SignatureScheme::ML_DSA_44,
164            Self::MlDsa65 => SignatureScheme::ML_DSA_65,
165            Self::MlDsa87 => SignatureScheme::ML_DSA_87,
166        }
167    }
168
169    fn alg_id(&self) -> AlgorithmIdentifier {
170        match self {
171            Self::MlDsa44 => alg_id::ML_DSA_44,
172            Self::MlDsa65 => alg_id::ML_DSA_65,
173            Self::MlDsa87 => alg_id::ML_DSA_87,
174        }
175    }
176}
177
178/// Keep in sync with the `SUPPORTED_SIG_ALGS` in `rustls_aws_lc_rs`.
179static SUPPORTED_SIG_ALGS: WebPkiSupportedAlgorithms = match WebPkiSupportedAlgorithms::new(
180    &[
181        rustls_aws_lc_rs::ECDSA_P256_SHA256,
182        rustls_aws_lc_rs::ECDSA_P256_SHA384,
183        rustls_aws_lc_rs::ECDSA_P384_SHA256,
184        rustls_aws_lc_rs::ECDSA_P384_SHA384,
185        rustls_aws_lc_rs::ECDSA_P521_SHA256,
186        rustls_aws_lc_rs::ECDSA_P521_SHA384,
187        rustls_aws_lc_rs::ECDSA_P521_SHA512,
188        rustls_aws_lc_rs::ED25519,
189        rustls_aws_lc_rs::RSA_PSS_2048_8192_SHA256_LEGACY_KEY,
190        rustls_aws_lc_rs::RSA_PSS_2048_8192_SHA384_LEGACY_KEY,
191        rustls_aws_lc_rs::RSA_PSS_2048_8192_SHA512_LEGACY_KEY,
192        rustls_aws_lc_rs::RSA_PKCS1_2048_8192_SHA256,
193        rustls_aws_lc_rs::RSA_PKCS1_2048_8192_SHA384,
194        rustls_aws_lc_rs::RSA_PKCS1_2048_8192_SHA512,
195        rustls_aws_lc_rs::RSA_PKCS1_2048_8192_SHA256_ABSENT_PARAMS,
196        rustls_aws_lc_rs::RSA_PKCS1_2048_8192_SHA384_ABSENT_PARAMS,
197        rustls_aws_lc_rs::RSA_PKCS1_2048_8192_SHA512_ABSENT_PARAMS,
198        ML_DSA_44,
199        ML_DSA_65,
200        ML_DSA_87,
201    ],
202    &[
203        // Note: for TLS1.2 the curve is not fixed by SignatureScheme. For TLS1.3 it is.
204        (
205            SignatureScheme::ECDSA_NISTP384_SHA384,
206            &[
207                rustls_aws_lc_rs::ECDSA_P384_SHA384,
208                rustls_aws_lc_rs::ECDSA_P256_SHA384,
209                rustls_aws_lc_rs::ECDSA_P521_SHA384,
210            ],
211        ),
212        (
213            SignatureScheme::ECDSA_NISTP256_SHA256,
214            &[
215                rustls_aws_lc_rs::ECDSA_P256_SHA256,
216                rustls_aws_lc_rs::ECDSA_P384_SHA256,
217                rustls_aws_lc_rs::ECDSA_P521_SHA256,
218            ],
219        ),
220        (
221            SignatureScheme::ECDSA_NISTP521_SHA512,
222            &[
223                rustls_aws_lc_rs::ECDSA_P521_SHA512,
224                rustls_aws_lc_rs::ECDSA_P384_SHA512,
225                rustls_aws_lc_rs::ECDSA_P256_SHA512,
226            ],
227        ),
228        (SignatureScheme::ED25519, &[rustls_aws_lc_rs::ED25519]),
229        (
230            SignatureScheme::RSA_PSS_SHA512,
231            &[rustls_aws_lc_rs::RSA_PSS_2048_8192_SHA512_LEGACY_KEY],
232        ),
233        (
234            SignatureScheme::RSA_PSS_SHA384,
235            &[rustls_aws_lc_rs::RSA_PSS_2048_8192_SHA384_LEGACY_KEY],
236        ),
237        (
238            SignatureScheme::RSA_PSS_SHA256,
239            &[rustls_aws_lc_rs::RSA_PSS_2048_8192_SHA256_LEGACY_KEY],
240        ),
241        (
242            SignatureScheme::RSA_PKCS1_SHA512,
243            &[rustls_aws_lc_rs::RSA_PKCS1_2048_8192_SHA512],
244        ),
245        (
246            SignatureScheme::RSA_PKCS1_SHA384,
247            &[rustls_aws_lc_rs::RSA_PKCS1_2048_8192_SHA384],
248        ),
249        (
250            SignatureScheme::RSA_PKCS1_SHA256,
251            &[rustls_aws_lc_rs::RSA_PKCS1_2048_8192_SHA256],
252        ),
253        (SignatureScheme::ML_DSA_44, &[ML_DSA_44]),
254        (SignatureScheme::ML_DSA_65, &[ML_DSA_65]),
255        (SignatureScheme::ML_DSA_87, &[ML_DSA_87]),
256    ],
257) {
258    Ok(algs) => algs,
259    Err(_) => panic!("bad WebPkiSupportedAlgorithms"),
260};
261
262/// ML-DSA signatures using the [4, 4] matrix (security strength category 2).
263pub static ML_DSA_44: &dyn SignatureVerificationAlgorithm = &AwsLcRsVerificationAlgorithm {
264    public_key_alg_id: alg_id::ML_DSA_44,
265    signature_alg_id: alg_id::ML_DSA_44,
266    verification_alg: &aws_lc_rs::unstable::signature::ML_DSA_44,
267    // Not included in AWS-LC-FIPS 3.0 FIPS scope
268    in_fips_submission: false,
269};
270
271/// ML-DSA signatures using the [6, 5] matrix (security strength category 3).
272pub static ML_DSA_65: &dyn SignatureVerificationAlgorithm = &AwsLcRsVerificationAlgorithm {
273    public_key_alg_id: alg_id::ML_DSA_65,
274    signature_alg_id: alg_id::ML_DSA_65,
275    verification_alg: &aws_lc_rs::unstable::signature::ML_DSA_65,
276    // Not included in AWS-LC-FIPS 3.0 FIPS scope
277    in_fips_submission: false,
278};
279
280/// ML-DSA signatures using the [8. 7] matrix (security strength category 5).
281pub static ML_DSA_87: &dyn SignatureVerificationAlgorithm = &AwsLcRsVerificationAlgorithm {
282    public_key_alg_id: alg_id::ML_DSA_87,
283    signature_alg_id: alg_id::ML_DSA_87,
284    verification_alg: &aws_lc_rs::unstable::signature::ML_DSA_87,
285    // Not included in AWS-LC-FIPS 3.0 FIPS scope
286    in_fips_submission: false,
287};
288
289#[cfg(test)]
290mod tests {
291    use rcgen::{
292        CertificateParams, CertifiedIssuer, ExtendedKeyUsagePurpose, IsCa, KeyPair, KeyUsagePurpose,
293    };
294    use rustls::crypto::Identity;
295    use rustls::{ClientConfig, RootCertStore, ServerConfig, ServerConnection, VecInput};
296    use rustls_test::do_handshake;
297
298    use super::*;
299
300    #[test]
301    fn ml_dsa() {
302        let ca_key = KeyPair::generate_for(&rcgen::PKCS_ML_DSA_44).unwrap();
303        let mut ca_params = CertificateParams::new(vec!["Test CA".into()]).unwrap();
304        ca_params.is_ca = IsCa::Ca(rcgen::BasicConstraints::Unconstrained);
305        ca_params.key_usages = vec![
306            KeyUsagePurpose::DigitalSignature,
307            KeyUsagePurpose::KeyCertSign,
308        ];
309        ca_params.extended_key_usages = vec![ExtendedKeyUsagePurpose::ServerAuth];
310        let issuer = CertifiedIssuer::self_signed(ca_params, ca_key).unwrap();
311
312        let ee_key = KeyPair::generate_for(&rcgen::PKCS_ML_DSA_87).unwrap();
313        let ee_params = CertificateParams::new(vec!["localhost".into()]).unwrap();
314        let ee_cert = ee_params
315            .signed_by(&ee_key, &issuer)
316            .unwrap();
317
318        let provider = Arc::new(DEFAULT_PROVIDER);
319        let server_config = ServerConfig::builder(provider.clone())
320            .with_no_client_auth()
321            .with_single_cert(
322                Arc::new(Identity::from_cert_chain(vec![ee_cert.der().clone()]).unwrap()),
323                PrivateKeyDer::try_from(ee_key.serialize_der()).unwrap(),
324            )
325            .unwrap();
326
327        let mut roots = RootCertStore::empty();
328        roots.add(issuer.der().clone()).unwrap();
329        let mut client = Arc::new(
330            ClientConfig::builder(provider)
331                .with_root_certificates(roots)
332                .with_no_client_auth()
333                .unwrap(),
334        )
335        .connect("localhost".try_into().unwrap())
336        .build()
337        .unwrap();
338
339        let mut client_input = VecInput::default();
340        let mut server_input = VecInput::default();
341        let mut server = ServerConnection::new(Arc::new(server_config)).unwrap();
342        do_handshake(
343            &mut client_input,
344            &mut client,
345            &mut server_input,
346            &mut server,
347        );
348    }
349}