1use std::fmt;
16
17use aws_lc_rs::aead::{Aad, LessSafeKey, Nonce, UnboundKey, AES_256_GCM};
18use aws_lc_rs::agreement::{self, ParsedPublicKey, UnparsedPublicKey, ECDH_P384};
19use aws_lc_rs::constant_time::verify_slices_are_equal;
20use aws_lc_rs::hmac;
21use aws_lc_rs::kem::{Ciphertext, DecapsulationKey, EncapsulationKey, ML_KEM_1024};
22use sha2::{Digest, Sha384};
23
24use crate::cbor::{self, Value};
25use crate::profile::Profile;
26
27mod keyring;
28pub use keyring::{Clock, Keyring, KEY_LIFETIME_MS, RETIRED_KEY_KEPT_MS};
29
30pub const SCHEME: i64 = 1;
32
33pub const KEY_ID_SIZE: usize = 8;
35pub const KEY_HASH_SIZE: usize = 48;
37pub const MLKEM_CIPHERTEXT_SIZE: usize = 1568;
39pub const MLKEM_DK_SIZE: usize = 3168;
41pub const NONCE_SIZE: usize = 12;
43pub const TAG_SIZE: usize = 16;
45
46const MLKEM_EK_BYTES: usize = 1568;
47const P384_POINT_BYTES: usize = 97;
48const P384_SCALAR_BYTES: usize = 48;
49
50pub const FRAME_CALL: &str = "call";
52pub const FRAME_STREAM_OPEN: &str = "stream_open";
53pub const FRAME_RESULT: &str = "result";
54pub const FRAME_ERROR: &str = "error";
55
56const LABEL_PURE: &str = "MACULA-E2E-PURE-V1";
57const LABEL_HYBRID: &str = "MACULA-E2E-HYBRID-V1";
58const LABEL_CALL: &str = "MACULA-E2E-CALL-V1";
59const LABEL_STREAM: &str = "MACULA-E2E-STREAM-V1";
60const LABEL_AAD: &str = "MACULA-E2E-AAD-V1";
61const LABEL_STREAM_AAD: &str = "MACULA-E2E-STREAM-AAD-V1";
62
63pub const MAX_ERROR_CODE_BYTES: usize = 64;
66pub const MAX_ERROR_DETAIL_BYTES: usize = 256;
67
68#[derive(Debug, Clone, PartialEq, Eq)]
70pub enum SealError {
71 Refused,
76 Key(String),
78 NotAnErrorPlain,
81 Unavailable,
83}
84
85impl fmt::Display for SealError {
86 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
87 match self {
88 SealError::Refused => f.write_str("sealed_refused"),
89 SealError::Key(why) => write!(f, "not a recipient key: {why}"),
90 SealError::NotAnErrorPlain => {
91 f.write_str("an ERROR's plaintext is not cbor([code, detail])")
92 }
93 SealError::Unavailable => f.write_str("a seal primitive is unavailable"),
94 }
95 }
96}
97
98impl std::error::Error for SealError {}
99
100pub fn carried_key_size(profile: Profile) -> usize {
102 match profile {
103 Profile::PqPure => MLKEM_EK_BYTES,
104 Profile::PqHybrid => MLKEM_EK_BYTES + P384_POINT_BYTES,
105 }
106}
107
108pub fn is_carried_key_size(len: usize) -> bool {
112 [Profile::PqPure, Profile::PqHybrid]
113 .into_iter()
114 .any(|p| carried_key_size(p) == len)
115}
116
117pub fn key_hash(carried: &[u8]) -> [u8; KEY_HASH_SIZE] {
119 Sha384::digest(carried).into()
120}
121
122pub fn key_id(carried: &[u8]) -> [u8; KEY_ID_SIZE] {
124 let mut id = [0; KEY_ID_SIZE];
125 id.copy_from_slice(&key_hash(carried)[..KEY_ID_SIZE]);
126 id
127}
128
129#[derive(Debug, Clone, PartialEq, Eq)]
131pub struct PublicKey {
132 profile: Profile,
133 carried: Vec<u8>,
134}
135
136impl PublicKey {
137 pub fn carried(&self) -> &[u8] {
139 &self.carried
140 }
141
142 pub fn profile(&self) -> Profile {
144 self.profile
145 }
146
147 pub fn key_id(&self) -> [u8; KEY_ID_SIZE] {
149 key_id(&self.carried)
150 }
151
152 fn mlkem(&self) -> &[u8] {
153 &self.carried[..MLKEM_EK_BYTES]
154 }
155
156 fn p384(&self) -> Option<&[u8]> {
157 (self.profile == Profile::PqHybrid).then(|| &self.carried[MLKEM_EK_BYTES..])
158 }
159}
160
161pub fn parse_public_key(profile: Profile, carried: &[u8]) -> Result<PublicKey, SealError> {
166 if carried.len() != carried_key_size(profile) {
167 return Err(SealError::Key(format!(
168 "a {} key as carried is {} bytes, not {}",
169 profile.name(),
170 carried_key_size(profile),
171 carried.len()
172 )));
173 }
174 EncapsulationKey::new(&ML_KEM_1024, &carried[..MLKEM_EK_BYTES])
175 .map_err(|e| SealError::Key(format!("ML-KEM-1024: {e}")))?;
176 if profile == Profile::PqHybrid && !p384_point(&carried[MLKEM_EK_BYTES..]) {
177 return Err(SealError::Key("not an uncompressed point on P-384".into()));
178 }
179 Ok(PublicKey {
180 profile,
181 carried: carried.to_vec(),
182 })
183}
184
185fn p384_point(point: &[u8]) -> bool {
187 point.len() == P384_POINT_BYTES
188 && point[0] == 0x04
189 && ParsedPublicKey::try_from(&UnparsedPublicKey::new(&ECDH_P384, point)).is_ok()
190}
191
192pub struct PrivateKey {
195 public: PublicKey,
196 mlkem: DecapsulationKey,
197 p384: Option<agreement::PrivateKey>,
198}
199
200impl fmt::Debug for PrivateKey {
201 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
202 write!(
203 f,
204 "seal::PrivateKey({}, key id {})",
205 self.public.profile.name(),
206 hex(&self.public.key_id())
207 )
208 }
209}
210
211impl PrivateKey {
212 pub fn generate(profile: Profile) -> Result<PrivateKey, SealError> {
214 let mlkem = DecapsulationKey::generate(&ML_KEM_1024).map_err(|_| SealError::Unavailable)?;
215 let ek = mlkem
216 .encapsulation_key()
217 .and_then(|k| k.key_bytes())
218 .map_err(|_| SealError::Unavailable)?;
219 let mut carried = ek.as_ref().to_vec();
220 let p384 = match profile {
221 Profile::PqPure => None,
222 Profile::PqHybrid => {
223 let key = agreement::PrivateKey::generate(&ECDH_P384)
224 .map_err(|_| SealError::Unavailable)?;
225 let point = key
226 .compute_public_key()
227 .map_err(|_| SealError::Unavailable)?;
228 carried.extend_from_slice(point.as_ref());
229 Some(key)
230 }
231 };
232 Ok(PrivateKey {
233 public: PublicKey { profile, carried },
234 mlkem,
235 p384,
236 })
237 }
238
239 pub fn from_parts(
243 profile: Profile,
244 mlkem_dk: &[u8],
245 carried: &[u8],
246 p384_scalar: Option<&[u8]>,
247 ) -> Result<PrivateKey, SealError> {
248 let public = parse_public_key(profile, carried)?;
249 if mlkem_dk.len() != MLKEM_DK_SIZE {
250 return Err(SealError::Key(format!(
251 "an ML-KEM-1024 decapsulation key is {MLKEM_DK_SIZE} bytes, not {}",
252 mlkem_dk.len()
253 )));
254 }
255 let mlkem = DecapsulationKey::new(&ML_KEM_1024, mlkem_dk)
256 .map_err(|e| SealError::Key(format!("ML-KEM-1024: {e}")))?;
257 let p384 = match (profile, p384_scalar) {
258 (Profile::PqPure, None) => None,
259 (Profile::PqHybrid, Some(scalar)) if scalar.len() == P384_SCALAR_BYTES => Some(
260 agreement::PrivateKey::from_private_key(&ECDH_P384, scalar)
261 .map_err(|e| SealError::Key(format!("P-384: {e}")))?,
262 ),
263 (profile, scalar) => {
264 return Err(SealError::Key(format!(
265 "a {} key with {} bytes of P-384 scalar",
266 profile.name(),
267 scalar.map_or(0, <[u8]>::len)
268 )))
269 }
270 };
271 if let (Some(key), Some(point)) = (&p384, public.p384()) {
272 let derived = key
273 .compute_public_key()
274 .map_err(|_| SealError::Unavailable)?;
275 if derived.as_ref() != point {
276 return Err(SealError::Key(
277 "the P-384 scalar is not the carried point's".into(),
278 ));
279 }
280 }
281 Ok(PrivateKey {
282 public,
283 mlkem,
284 p384,
285 })
286 }
287
288 pub fn public_key(&self) -> &PublicKey {
290 &self.public
291 }
292
293 pub fn mlkem_dk(&self) -> Result<Vec<u8>, SealError> {
296 self.mlkem
297 .key_bytes()
298 .map(|b| b.as_ref().to_vec())
299 .map_err(|_| SealError::Unavailable)
300 }
301}
302
303pub fn sender_secret(recipient: &PublicKey) -> Result<([u8; KEY_HASH_SIZE], Vec<u8>), SealError> {
307 let ek = EncapsulationKey::new(&ML_KEM_1024, recipient.mlkem())
308 .map_err(|_| SealError::Unavailable)?;
309 let (mlkem_ct, ss_mlkem) = ek.encapsulate().map_err(|_| SealError::Unavailable)?;
310 let mlkem_ct = mlkem_ct.as_ref().to_vec();
311 let ss_mlkem = ss_mlkem.as_ref().to_vec();
312 let key_hash = key_hash(recipient.carried());
313 let Some(point) = recipient.p384() else {
314 return Ok((
315 combine(LABEL_PURE, &pure_ikm(&ss_mlkem, &mlkem_ct, &key_hash)),
316 mlkem_ct,
317 ));
318 };
319 let ephemeral =
320 agreement::PrivateKey::generate(&ECDH_P384).map_err(|_| SealError::Unavailable)?;
321 let eph_pub = ephemeral
322 .compute_public_key()
323 .map_err(|_| SealError::Unavailable)?
324 .as_ref()
325 .to_vec();
326 let ss_ecdh = ecdh_secret(&ephemeral, point)?;
327 let ss = combine(
328 LABEL_HYBRID,
329 &hybrid_ikm(&ss_mlkem, &ss_ecdh, &mlkem_ct, &eph_pub, &key_hash),
330 );
331 Ok((ss, [mlkem_ct, eph_pub].concat()))
332}
333
334pub fn recipient_secret(
339 recipient: &PrivateKey,
340 kem_ct: &[u8],
341) -> Result<[u8; KEY_HASH_SIZE], SealError> {
342 let key_hash = key_hash(recipient.public.carried());
343 let decapsulated = |ct: &[u8]| {
344 recipient
345 .mlkem
346 .decapsulate(Ciphertext::from(ct))
347 .map(|ss| ss.as_ref().to_vec())
348 .map_err(|_| SealError::Refused)
349 };
350 let Some(p384) = &recipient.p384 else {
351 if kem_ct.len() != MLKEM_CIPHERTEXT_SIZE {
352 return Err(SealError::Refused);
353 }
354 let ss_mlkem = decapsulated(kem_ct)?;
355 return Ok(combine(LABEL_PURE, &pure_ikm(&ss_mlkem, kem_ct, &key_hash)));
356 };
357 if kem_ct.len() != MLKEM_CIPHERTEXT_SIZE + P384_POINT_BYTES {
358 return Err(SealError::Refused);
359 }
360 let (mlkem_ct, eph_pub) = kem_ct.split_at(MLKEM_CIPHERTEXT_SIZE);
361 let ss_ecdh = ecdh_secret(p384, eph_pub)?;
362 let ss_mlkem = decapsulated(mlkem_ct)?;
363 Ok(combine(
364 LABEL_HYBRID,
365 &hybrid_ikm(&ss_mlkem, &ss_ecdh, mlkem_ct, eph_pub, &key_hash),
366 ))
367}
368
369fn ecdh_secret(private: &agreement::PrivateKey, peer_point: &[u8]) -> Result<Vec<u8>, SealError> {
374 if !p384_point(peer_point) {
375 return Err(SealError::Refused);
376 }
377 let secret = agreement::agree(
378 private,
379 UnparsedPublicKey::new(&ECDH_P384, peer_point),
380 SealError::Refused,
381 |secret| Ok(secret.to_vec()),
382 )?;
383 let zero = verify_slices_are_equal(&secret, &[0; P384_SCALAR_BYTES]).is_ok();
384 if secret.len() != P384_SCALAR_BYTES || zero {
385 return Err(SealError::Refused);
386 }
387 Ok(secret)
388}
389
390fn pure_ikm(ss_mlkem: &[u8], mlkem_ct: &[u8], key_hash: &[u8]) -> Vec<u8> {
391 encode(vec![bytes(ss_mlkem), bytes(mlkem_ct), bytes(key_hash)])
392}
393
394fn hybrid_ikm(
395 ss_mlkem: &[u8],
396 ss_ecdh: &[u8],
397 mlkem_ct: &[u8],
398 eph_pub: &[u8],
399 key_hash: &[u8],
400) -> Vec<u8> {
401 encode(vec![
402 bytes(ss_mlkem),
403 bytes(ss_ecdh),
404 bytes(mlkem_ct),
405 bytes(eph_pub),
406 bytes(key_hash),
407 ])
408}
409
410fn combine(label: &str, ikm: &[u8]) -> [u8; KEY_HASH_SIZE] {
412 extract(label.as_bytes(), ikm)
413}
414
415#[derive(Debug, Clone, Copy, PartialEq, Eq)]
418pub struct Parties {
419 pub request_id: [u8; 16],
420 pub caller: [u8; 32],
421 pub target: [u8; 32],
422}
423
424pub fn call_keys(
427 secret: &[u8; KEY_HASH_SIZE],
428 frame_type: &str,
429 p: &Parties,
430) -> ([u8; 32], [u8; 32]) {
431 halves(expand(
432 secret,
433 &encode(vec![
434 Value::text(LABEL_CALL),
435 Value::text(frame_type),
436 bytes(&p.request_id),
437 bytes(&p.caller),
438 bytes(&p.target),
439 ]),
440 ))
441}
442
443pub fn stream_keys(secret: &[u8; KEY_HASH_SIZE], p: &Parties) -> ([u8; 32], [u8; 32]) {
445 halves(expand(
446 secret,
447 &encode(vec![
448 Value::text(LABEL_STREAM),
449 bytes(&p.request_id),
450 bytes(&p.caller),
451 bytes(&p.target),
452 ]),
453 ))
454}
455
456fn halves(okm: [u8; 64]) -> ([u8; 32], [u8; 32]) {
457 let mut a = [0; 32];
458 let mut b = [0; 32];
459 a.copy_from_slice(&okm[..32]);
460 b.copy_from_slice(&okm[32..]);
461 (a, b)
462}
463
464#[derive(Debug, Clone, PartialEq, Eq)]
466pub struct Request {
467 pub frame_type: String,
468 pub realm: [u8; 32],
469 pub procedure: String,
470 pub caller: [u8; 32],
471 pub target: [u8; 32],
472 pub request_id: [u8; 16],
473 pub deadline: u64,
474}
475
476impl Request {
477 fn fields(&self, frame_type: &str) -> Vec<Value> {
478 vec![
479 Value::text(LABEL_AAD),
480 Value::text(frame_type),
481 bytes(&self.realm),
482 Value::text(self.procedure.clone()),
483 bytes(&self.caller),
484 bytes(&self.target),
485 bytes(&self.request_id),
486 Value::Int(i128::from(self.deadline)),
487 ]
488 }
489}
490
491pub fn request_aad(r: &Request) -> Vec<u8> {
494 encode(r.fields(&r.frame_type))
495}
496
497pub fn reply_aad(
501 r: &Request,
502 reply_frame_type: &str,
503 request_hash: &[u8; 48],
504 responded_by: &[u8; 32],
505) -> Vec<u8> {
506 let mut fields = r.fields(reply_frame_type);
507 fields.push(bytes(request_hash));
508 fields.push(bytes(responded_by));
509 encode(fields)
510}
511
512#[derive(Debug, Clone, Copy, PartialEq, Eq)]
514pub enum Direction {
515 CallerToProvider = 0,
517 ProviderToCaller = 1,
519}
520
521pub fn stream_aad(frame_type: &str, request_id: &[u8; 16], seq: u64, d: Direction) -> Vec<u8> {
523 encode(vec![
524 Value::text(LABEL_STREAM_AAD),
525 Value::text(frame_type),
526 bytes(request_id),
527 Value::Int(i128::from(seq)),
528 Value::Int(d as i128),
529 ])
530}
531
532pub fn stream_nonce(seq: u64) -> [u8; NONCE_SIZE] {
534 let mut nonce = [0; NONCE_SIZE];
535 nonce[4..].copy_from_slice(&seq.to_be_bytes());
536 nonce
537}
538
539pub fn random_nonce() -> Result<[u8; NONCE_SIZE], SealError> {
542 let mut nonce = [0; NONCE_SIZE];
543 aws_lc_rs::rand::fill(&mut nonce).map_err(|_| SealError::Unavailable)?;
544 Ok(nonce)
545}
546
547pub fn seal(key: &[u8; 32], nonce: &[u8; NONCE_SIZE], aad: &[u8], plain: &[u8]) -> Vec<u8> {
549 let key = LessSafeKey::new(UnboundKey::new(&AES_256_GCM, key).expect("a 32-byte AES-256 key"));
550 let mut out = plain.to_vec();
551 key.seal_in_place_append_tag(
552 Nonce::assume_unique_for_key(*nonce),
553 Aad::from(aad),
554 &mut out,
555 )
556 .expect("AES-256-GCM seals any payload under 2^36 bytes");
557 out
558}
559
560pub fn open(
563 key: &[u8; 32],
564 nonce: &[u8; NONCE_SIZE],
565 aad: &[u8],
566 sealed: &[u8],
567) -> Result<Vec<u8>, SealError> {
568 if sealed.len() < TAG_SIZE {
569 return Err(SealError::Refused);
570 }
571 let key = LessSafeKey::new(UnboundKey::new(&AES_256_GCM, key).expect("a 32-byte AES-256 key"));
572 let mut buf = sealed.to_vec();
573 let plain = key
574 .open_in_place(
575 Nonce::assume_unique_for_key(*nonce),
576 Aad::from(aad),
577 &mut buf,
578 )
579 .map_err(|_| SealError::Refused)?;
580 Ok(plain.to_vec())
581}
582
583pub fn error_plain(code: &str, detail: &str) -> Vec<u8> {
587 encode(vec![Value::text(code), Value::text(detail)])
588}
589
590pub fn open_error_plain(plain: &[u8]) -> Result<(String, Option<String>), SealError> {
593 match cbor::decode(plain) {
594 Ok(Value::List(items)) => match items.as_slice() {
595 [Value::Text(code), Value::Text(detail)]
596 if code.len() <= MAX_ERROR_CODE_BYTES && detail.len() <= MAX_ERROR_DETAIL_BYTES =>
597 {
598 Ok((code.clone(), (!detail.is_empty()).then(|| detail.clone())))
599 }
600 _ => Err(SealError::NotAnErrorPlain),
601 },
602 _ => Err(SealError::NotAnErrorPlain),
603 }
604}
605
606fn extract(salt: &[u8], ikm: &[u8]) -> [u8; KEY_HASH_SIZE] {
610 let tag = hmac::sign(&hmac::Key::new(hmac::HMAC_SHA384, salt), ikm);
611 let mut prk = [0; KEY_HASH_SIZE];
612 prk.copy_from_slice(tag.as_ref());
613 prk
614}
615
616fn expand(prk: &[u8; KEY_HASH_SIZE], info: &[u8]) -> [u8; 64] {
619 let key = hmac::Key::new(hmac::HMAC_SHA384, prk);
620 let mut okm = [0; 64];
621 let mut previous: Vec<u8> = Vec::new();
622 for (counter, chunk) in (1u8..).zip(okm.chunks_mut(KEY_HASH_SIZE)) {
623 let mut ctx = hmac::Context::with_key(&key);
624 ctx.update(&previous);
625 ctx.update(info);
626 ctx.update(&[counter]);
627 previous = ctx.sign().as_ref().to_vec();
628 chunk.copy_from_slice(&previous[..chunk.len()]);
629 }
630 okm
631}
632
633fn bytes(b: &[u8]) -> Value {
634 Value::Bytes(b.to_vec())
635}
636
637fn encode(items: Vec<Value>) -> Vec<u8> {
641 cbor::encode(&Value::List(items)).expect("seal arrays hold no integer outside u64")
642}
643
644fn hex(b: &[u8]) -> String {
645 b.iter().map(|x| format!("{x:02x}")).collect()
646}
647
648#[cfg(test)]
649mod vectors;