1use 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
31pub 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 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
178static 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 (
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
262pub 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 in_fips_submission: false,
269};
270
271pub 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 in_fips_submission: false,
278};
279
280pub 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 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}