1use aws_lc_rs::signature::{
33 KeyPair as _, ML_DSA_65, ML_DSA_65_SIGNING, ParsedPublicKey, PqdsaKeyPair,
34};
35use base64ct::{Base64UrlUnpadded, Encoding as _};
36use sha2::{Digest as _, Sha256, Sha512};
37
38use crate::misc::error::{OPAQUE, Opaque};
39use crate::misc::jwt;
40use crate::misc::serde_ext::bytes_wrapper::B64;
41
42pub const ALG: &str = "ph-ML-DSA-65-Ed25519";
45
46const PREFIX: &[u8] = b"CompositeAlgorithmSignatures2025";
49
50const LABEL: &[u8] = b"COMPSIG-MLDSA65-Ed25519-SHA512";
56
57const CTX: &[u8] = b"";
60
61const _: () = assert!(CTX.len() <= 255);
63
64const ML_DSA_CTX: &[u8] = b"";
70
71const _: () = assert!(ML_DSA_CTX.is_empty());
76
77const ED25519_SIG_LEN: usize = ed25519_dalek::SIGNATURE_LENGTH;
79
80const M_PRIME_LEN: usize = PREFIX.len() + LABEL.len() + 1 + CTX.len() + 64;
83
84#[derive(Debug)]
86pub struct SigningKey {
87 ed: ed25519_dalek::SigningKey,
88 ml: PqdsaKeyPair,
89
90 ml_seed: zeroize::Zeroizing<[u8; 32]>,
93
94 vk: VerifyingKey,
97}
98
99#[derive(Clone, Debug)]
101pub struct VerifyingKey {
102 ed: ed25519_dalek::VerifyingKey,
103
104 ml: ParsedPublicKey,
106}
107
108#[derive(Clone, Debug, serde::Serialize, serde::Deserialize, zeroize::ZeroizeOnDrop)]
110pub struct SigningKeyBytes {
111 ed: B64,
113
114 ml: B64,
116}
117
118#[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
120pub struct VerifyingKeyBytes {
121 pub ed: B64,
123
124 pub ml: B64,
126}
127
128impl SigningKey {
129 pub fn generate() -> Result<Self, Opaque> {
131 Self::from_parts(
132 crate::misc::crypto::random_32_bytes(),
133 crate::misc::crypto::random_32_bytes(),
134 )
135 }
136
137 pub fn verifying_key(&self) -> &VerifyingKey {
142 &self.vk
143 }
144
145 pub fn encode(&self) -> SigningKeyBytes {
147 SigningKeyBytes {
148 ed: B64::from_bytes(self.ed.to_bytes()),
149 ml: B64::from_bytes(*self.ml_seed),
150 }
151 }
152
153 pub fn ed25519_signing_key(&self) -> &ed25519_dalek::SigningKey {
159 &self.ed
160 }
161
162 fn from_parts(ed_seed: [u8; 32], ml_seed: [u8; 32]) -> Result<Self, Opaque> {
167 let ed = ed25519_dalek::SigningKey::from_bytes(&ed_seed);
168 let ml = PqdsaKeyPair::from_seed(&ML_DSA_65_SIGNING, &ml_seed).map_err(|_| OPAQUE)?;
169 let vk = VerifyingKey {
170 ed: ed.verifying_key(),
171 ml: ParsedPublicKey::new(&ML_DSA_65, ml.public_key().as_ref()).map_err(|_| OPAQUE)?,
172 };
173 Ok(Self {
174 ed,
175 ml,
176 ml_seed: zeroize::Zeroizing::new(ml_seed),
177 vk,
178 })
179 }
180}
181
182impl SigningKeyBytes {
183 pub fn decode(&self) -> Result<SigningKey, Opaque> {
185 let ed_seed: [u8; 32] = (&self.ed[..]).try_into()?;
186 let ml_seed: [u8; 32] = (&self.ml[..]).try_into()?;
187 SigningKey::from_parts(ed_seed, ml_seed)
188 }
189}
190
191impl VerifyingKey {
192 pub fn encode(&self) -> VerifyingKeyBytes {
194 VerifyingKeyBytes {
195 ed: B64::from_bytes(self.ed.to_bytes()),
196 ml: B64::from_bytes(self.ml.as_ref()),
197 }
198 }
199
200 pub fn ed25519_bytes(&self) -> [u8; 32] {
203 self.ed.to_bytes()
204 }
205}
206
207impl VerifyingKeyBytes {
208 pub fn decode(&self) -> Result<VerifyingKey, Opaque> {
210 let ed_bytes: [u8; 32] = (&self.ed[..]).try_into()?;
211 Ok(VerifyingKey {
212 ed: ed25519_dalek::VerifyingKey::from_bytes(&ed_bytes)?,
213 ml: ParsedPublicKey::new(&ML_DSA_65, &self.ml[..]).map_err(|_| OPAQUE)?,
214 })
215 }
216}
217
218impl PartialEq for VerifyingKey {
221 fn eq(&self, other: &Self) -> bool {
222 self.ed == other.ed && self.ml.as_ref() == other.ml.as_ref()
223 }
224}
225
226impl Eq for VerifyingKey {}
227
228impl jwt::Key for SigningKey {
229 const ALG: &'static str = ALG;
230}
231
232impl jwt::Key for VerifyingKey {
233 const ALG: &'static str = ALG;
234}
235
236impl jwt::SigningKey for SigningKey {
237 type Signature = Vec<u8>;
238
239 fn sign(&self, message: &[u8]) -> anyhow::Result<Vec<u8>> {
240 let m_prime = message_representative(message);
241
242 let ml_sig_len = ML_DSA_65_SIGNING.signature_len();
244 let mut signature = vec![0u8; ml_sig_len + ED25519_SIG_LEN];
245 ml_dsa_sign(&self.ml, &m_prime, &mut signature[..ml_sig_len])
246 .map_err(|_| anyhow::anyhow!("ML-DSA signing failed"))?;
247 signature[ml_sig_len..]
248 .copy_from_slice(&ed25519_dalek::Signer::sign(&self.ed, &m_prime).to_bytes());
249 Ok(signature)
250 }
251
252 fn jwk(&self) -> serde_json::Value {
253 let mut pk = self.ml.public_key().as_ref().to_vec();
255 pk.extend_from_slice(self.ed.verifying_key().as_bytes());
256
257 serde_json::json!({
258 "kty": "AKP",
259 "alg": ALG,
260 "pub": Base64UrlUnpadded::encode_string(&pk),
261 "kid": jwk_thumbprint(ALG, &pk),
262 "use": "sig",
263 })
264 }
265}
266
267impl jwt::VerifyingKey for VerifyingKey {
268 fn is_valid_signature(&self, message: &[u8], signature: Vec<u8>) -> bool {
269 let ml_sig_len = ML_DSA_65_SIGNING.signature_len();
270 if signature.len() != ml_sig_len + ED25519_SIG_LEN {
271 return false;
272 }
273 let (ml_sig, ed_sig) = signature.split_at(ml_sig_len);
274
275 let m_prime = message_representative(message);
276
277 if !ml_dsa_verify(&self.ml, &m_prime, ml_sig) {
279 return false;
280 }
281 let Ok(ed_sig) = ed25519_dalek::Signature::from_slice(ed_sig) else {
282 return false;
283 };
284 ed25519_dalek::Verifier::verify(&self.ed, &m_prime, &ed_sig).is_ok()
285 }
286
287 fn describe(&self) -> String {
288 let mut pubkey = self.ml.as_ref().to_vec();
291 pubkey.extend_from_slice(self.ed.as_bytes());
292 format!("{ALG} key #{}", jwk_thumbprint(ALG, &pubkey))
293 }
294}
295
296fn jwk_thumbprint(alg: &str, pub_bytes: &[u8]) -> String {
303 let canonical = format!(
304 r#"{{"alg":"{alg}","kty":"AKP","pub":"{}"}}"#,
305 Base64UrlUnpadded::encode_string(pub_bytes)
306 );
307 Base64UrlUnpadded::encode_string(Sha256::digest(canonical.as_bytes()).as_slice())
308}
309
310fn message_representative(message: &[u8]) -> [u8; M_PRIME_LEN] {
316 let mut m_prime = [0u8; M_PRIME_LEN];
317 let mut at = 0;
318 m_prime[at..at + PREFIX.len()].copy_from_slice(PREFIX);
319 at += PREFIX.len();
320 m_prime[at..at + LABEL.len()].copy_from_slice(LABEL);
321 at += LABEL.len();
322 m_prime[at] = CTX.len() as u8; at += 1;
324 m_prime[at..at + CTX.len()].copy_from_slice(CTX);
325 at += CTX.len();
326 m_prime[at..].copy_from_slice(Sha512::digest(message).as_slice());
327 m_prime
328}
329
330fn ml_dsa_sign(keypair: &PqdsaKeyPair, m_prime: &[u8], out: &mut [u8]) -> Result<(), Opaque> {
337 keypair.sign(m_prime, out).map_err(|_| OPAQUE)?;
339 Ok(())
340}
341
342fn ml_dsa_verify(public_key: &ParsedPublicKey, m_prime: &[u8], signature: &[u8]) -> bool {
344 public_key.verify_sig(m_prime, signature).is_ok()
345}
346
347#[cfg(test)]
348mod tests {
349 use super::*;
350 use crate::misc::jwt::{self, Claims, SigningKey as _, VerifyingKey as _};
351 use base64ct::Base64UrlUnpadded;
352
353 #[test]
358 fn jwk_thumbprint_matches_rfc9964() {
359 const PUB: &str = "unH59k4RuutY-pxvu24U5h8YZD2rSVtHU5qRZsoBmBMcRPgmu9VuNOVdteXi1zNIXjnqJg_GAAxepLqA00Vc3lO0bzRIKu39VFD8Lhuk8l0V-cFEJC-zm7UihxiQMMUEmOFxe3x1ixkKZ0jqmqP3rKryx8tSbtcXyfea64QhT6XNje2SoMP6FViBDxLHBQo2dwjRls0k5a-XSQSu2OTOiHLoaWsLe8pQ5FLNfTDqmkrawDEdZyxr3oSWJAsHQxRjcIiVzZuvwxYy1zl2STiP2vy_fTBaPemkleynQzqPg7oPCyXEE8bjnJbrfWkbNNN8438e6tHPIX4l7zTuzz98YPhLjt_d6EBdT4MldsYe-Y4KLyjaGHcAlTkk9oa5RhRwW89T0z_t1DSO3dvfKLUGXh8gd1BD6Fz5MfgpF5NjoafnQEqDjsAAhrCXY4b-Y3yYJEdX4_dp3dRGdHG_rWcPmgX4JG7lCnser4f8QGnDriqiAzJYEXeS8LzUngg_0bx0lqv_KcyU5IaLISFO0xZSU5mmEPvdSoDnyAcV8pV44qhLtAvd29n0ehG259oRihtljTWeiu9V60a1N2tbZVl5mEqSK-6_xZvNYA1TCdzNctvweH24unV7U3wer9XA9Q6kvJWDVJ4oKaQsKMrCSMlteBJMRxWbGK7ddUq6F7GdQw-3j2M-qdJvVKm9UPjY9rc1lPgol25-oJxTu7nxGlbJUH-4m5pevAN6NyZ6lfhbjWTKlxkrEKZvQXs_Yf6cpXEwpI_ZJeriq1UC1XHIpRkDwdOY9MH3an4RdDl2r9vGl_IwlKPNdh_5aF3jLgn7PCit1FNJAwC8fIncAXgAlgcXIpRXdfJk4bBiO89GGccSyDh2EgXYdpG3XvNgGWy7npuSoNTE7WIyblAk13UQuO4sdCbMIuriCdyfE73mvwj15xgb07RZRQtFGlFTmnFcIdZ90zDrWXDbANntv7KCKwNvoTuv64bY3HiGbj-NQ-U9eMylWVpvr4hrXcES8c9K3PqHWADZC0iIOvlzFv4VBoc_wVflcOrL_SIoaNFCNBAZZq-2v5lAgpJTqVOtqJ_HVraoSfcKy5g45p-qULunXj6Jwq21fobQiKubBKKOZwcJFyJD7F4ACKXOrz-HIvSHMCWW_9dVrRuCpJw0s0aVFbRqopDNhu446nqb4_EDYQM1tTHMozPd_jKxRRD0sH75X8ZoToxFSpLBDbtdWcenxj-zBf6IGWfZnmaetjKEBYJWC7QDQx1A91pJVJCEgieCkoIfTqkeQuePpIyu48g2FG3P1zjRF-kumhUTfSjo5qS0YiZQy0E1BMs6M11EvuxXRsHClLHoy5nLYI2Sj4zjVjYyxSHyPRPGGo9hwB34yWxzYNtPPGiqXS_dNCpi_zRZwRY4lCGrQ-hYTEWIK1Dm5OlttvC4_eiQ1dv63NiGkLRJ5kJA3bICN0fzCDY-MBqnd1cWn8YVBijVkgtaoascjL9EywDgJdeHnXK0eeOvUxHHhXJVkNqcibn8O4RQdpVU60TSA-uiu675ytIjcBHC6kTv8A8pmkj_4oypPd-F92YIJC741swkYQoeIHj8rE-ThcMUkF7KqC5VORbZTRp8HsZSqgiJcIPaouuxd1-8Rxrid3fXkE6p8bkrysPYoxWEJgh7ZFsRCPDWX-yTeJwFN0PKFP1j0F6YtlLfK5wv-c4F8ZQHA_-yc_gODicy7KmWDZgbTP07e7gEWzw4MFRrndjbDQ";
360 const KID: &str = "T4xl70S7MT6Zeq6r9V9fPJGVn76wfnXJ21-gyo0Gu6o";
361
362 let pub_bytes = Base64UrlUnpadded::decode_vec(PUB).unwrap();
363 assert_eq!(jwk_thumbprint("ML-DSA-44", &pub_bytes), KID);
364 }
365
366 #[test]
373 fn jwk_kid_matches_describe() {
374 let sk = SigningKey::generate().unwrap();
375 let vk = sk.verifying_key();
376
377 let jwk_kid = sk.jwk()["kid"].as_str().unwrap().to_string();
378 let describe_kid = vk.describe().rsplit_once('#').unwrap().1.to_string();
380
381 assert_eq!(jwk_kid, describe_kid);
382 }
383
384 #[test]
385 fn sign_verify_and_tamper() {
386 let sk = SigningKey::generate().unwrap();
387 let vk = sk.verifying_key();
388 let message = b"the message to be signed";
389
390 let signature = sk.sign(message).unwrap();
391 assert!(vk.is_valid_signature(message, signature.clone()));
392 assert!(!vk.is_valid_signature(b"other message", signature.clone()));
394
395 let mut ml_tampered = signature.clone();
397 ml_tampered[0] ^= 1;
398 assert!(!vk.is_valid_signature(message, ml_tampered));
399
400 let mut ed_tampered = signature;
402 *ed_tampered.last_mut().unwrap() ^= 1;
403 assert!(!vk.is_valid_signature(message, ed_tampered));
404 }
405
406 #[test]
407 fn malformed_signature_rejected() {
408 let sk = SigningKey::generate().unwrap();
409 let vk = sk.verifying_key();
410 let message = b"msg";
411 let valid = sk.sign(message).unwrap();
412
413 assert!(!vk.is_valid_signature(message, vec![]));
415 assert!(!vk.is_valid_signature(message, valid[..valid.len() - 1].to_vec()));
416 let mut too_long = valid.clone();
417 too_long.push(0);
418 assert!(!vk.is_valid_signature(message, too_long));
419 }
420
421 #[test]
422 fn jwt_roundtrip() {
423 let sk = SigningKey::generate().unwrap();
424 let vk = sk.verifying_key();
425
426 let token = Claims::new()
427 .claim("foo", "bar")
428 .unwrap()
429 .sign(&sk)
430 .unwrap();
431 let mut claims = token.open(vk).unwrap();
432 assert_eq!(
433 claims.extract::<String>("foo").unwrap(),
434 Some("bar".to_string())
435 );
436
437 let hs_token = Claims::new().sign(&jwt::HS256(vec![0u8; 32])).unwrap();
439 assert!(matches!(
440 hs_token.open(vk),
441 Err(jwt::Error::UnexpectedAlgorithm { .. })
442 ));
443 }
444
445 #[test]
446 fn signed_roundtrip() {
447 #[derive(serde::Serialize, serde::Deserialize, PartialEq, Eq, Debug)]
448 struct TestMsg {
449 hello: String,
450 }
451 crate::api::having_message_code! { TestMsg, Example }
452
453 let sk = SigningKey::generate().unwrap();
454 let vk = sk.verifying_key();
455 let message = TestMsg {
456 hello: "world".to_string(),
457 };
458
459 let signed =
460 crate::api::Signed::<TestMsg>::new(&sk, &message, std::time::Duration::from_secs(60))
461 .unwrap();
462 assert_eq!(signed.open(vk, None).unwrap(), message);
463 }
464
465 #[test]
466 fn encode_decode_roundtrip() {
467 let sk = SigningKey::generate().unwrap();
468 let vk = sk.verifying_key();
469
470 let skb: SigningKeyBytes =
472 serde_json::from_str(&serde_json::to_string(&sk.encode()).unwrap()).unwrap();
473 let sk2 = skb.decode().unwrap();
474 assert_eq!(sk2.verifying_key(), vk);
475
476 let vkb: VerifyingKeyBytes =
478 serde_json::from_str(&serde_json::to_string(&vk.encode()).unwrap()).unwrap();
479 assert_eq!(vkb, vk.encode());
480 let vk2 = vkb.decode().unwrap();
481 assert_eq!(&vk2, vk);
482 }
483
484 #[test]
492 #[ignore = "regenerates the Python hub fixture; see the doc comment"]
493 fn emit_bespoke_test_vector() {
494 let sk = SigningKey::from_parts([0x11u8; 32], [0x22u8; 32]).unwrap();
495 let vk = sk.verifying_key();
496
497 let jws: String = Claims::new()
498 .claim("msg", "bespoke composite test vector")
499 .unwrap()
500 .sign(&sk)
501 .unwrap()
502 .into();
503
504 let vector = serde_json::json!({
505 "alg": ALG,
506 "verifying_key": serde_json::to_value(vk.encode()).unwrap(),
507 "jws": jws,
508 });
509 println!(
510 "BESPOKE_VECTOR_BEGIN\n{}\nBESPOKE_VECTOR_END",
511 serde_json::to_string_pretty(&vector).unwrap()
512 );
513 }
514}