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 check_scalar_answers_point(key, point)?;
273 }
274 Ok(PrivateKey {
275 public,
276 mlkem,
277 p384,
278 })
279 }
280
281 pub fn public_key(&self) -> &PublicKey {
283 &self.public
284 }
285
286 pub fn mlkem_dk(&self) -> Result<Vec<u8>, SealError> {
289 self.mlkem
290 .key_bytes()
291 .map(|b| b.as_ref().to_vec())
292 .map_err(|_| SealError::Unavailable)
293 }
294}
295
296fn check_scalar_answers_point(key: &agreement::PrivateKey, point: &[u8]) -> Result<(), SealError> {
298 let derived = key
299 .compute_public_key()
300 .map_err(|_| SealError::Unavailable)?;
301 if derived.as_ref() != point {
302 return Err(SealError::Key(
303 "the P-384 scalar is not the carried point's".into(),
304 ));
305 }
306 Ok(())
307}
308
309pub fn sender_secret(recipient: &PublicKey) -> Result<([u8; KEY_HASH_SIZE], Vec<u8>), SealError> {
313 let ek = EncapsulationKey::new(&ML_KEM_1024, recipient.mlkem())
314 .map_err(|_| SealError::Unavailable)?;
315 let (mlkem_ct, ss_mlkem) = ek.encapsulate().map_err(|_| SealError::Unavailable)?;
316 let mlkem_ct = mlkem_ct.as_ref().to_vec();
317 let ss_mlkem = ss_mlkem.as_ref().to_vec();
318 let key_hash = key_hash(recipient.carried());
319 let Some(point) = recipient.p384() else {
320 return Ok((
321 combine(LABEL_PURE, &pure_ikm(&ss_mlkem, &mlkem_ct, &key_hash)),
322 mlkem_ct,
323 ));
324 };
325 let ephemeral =
326 agreement::PrivateKey::generate(&ECDH_P384).map_err(|_| SealError::Unavailable)?;
327 let eph_pub = ephemeral
328 .compute_public_key()
329 .map_err(|_| SealError::Unavailable)?
330 .as_ref()
331 .to_vec();
332 let ss_ecdh = ecdh_secret(&ephemeral, point)?;
333 let ss = combine(
334 LABEL_HYBRID,
335 &hybrid_ikm(&ss_mlkem, &ss_ecdh, &mlkem_ct, &eph_pub, &key_hash),
336 );
337 Ok((ss, [mlkem_ct, eph_pub].concat()))
338}
339
340pub fn recipient_secret(
345 recipient: &PrivateKey,
346 kem_ct: &[u8],
347) -> Result<[u8; KEY_HASH_SIZE], SealError> {
348 let key_hash = key_hash(recipient.public.carried());
349 let decapsulated = |ct: &[u8]| {
350 recipient
351 .mlkem
352 .decapsulate(Ciphertext::from(ct))
353 .map(|ss| ss.as_ref().to_vec())
354 .map_err(|_| SealError::Refused)
355 };
356 let Some(p384) = &recipient.p384 else {
357 if kem_ct.len() != MLKEM_CIPHERTEXT_SIZE {
358 return Err(SealError::Refused);
359 }
360 let ss_mlkem = decapsulated(kem_ct)?;
361 return Ok(combine(LABEL_PURE, &pure_ikm(&ss_mlkem, kem_ct, &key_hash)));
362 };
363 if kem_ct.len() != MLKEM_CIPHERTEXT_SIZE + P384_POINT_BYTES {
364 return Err(SealError::Refused);
365 }
366 let (mlkem_ct, eph_pub) = kem_ct.split_at(MLKEM_CIPHERTEXT_SIZE);
367 let ss_ecdh = ecdh_secret(p384, eph_pub)?;
368 let ss_mlkem = decapsulated(mlkem_ct)?;
369 Ok(combine(
370 LABEL_HYBRID,
371 &hybrid_ikm(&ss_mlkem, &ss_ecdh, mlkem_ct, eph_pub, &key_hash),
372 ))
373}
374
375fn ecdh_secret(private: &agreement::PrivateKey, peer_point: &[u8]) -> Result<Vec<u8>, SealError> {
380 if !p384_point(peer_point) {
381 return Err(SealError::Refused);
382 }
383 let secret = agreement::agree(
384 private,
385 UnparsedPublicKey::new(&ECDH_P384, peer_point),
386 SealError::Refused,
387 |secret| Ok(secret.to_vec()),
388 )?;
389 let zero = verify_slices_are_equal(&secret, &[0; P384_SCALAR_BYTES]).is_ok();
390 if secret.len() != P384_SCALAR_BYTES || zero {
391 return Err(SealError::Refused);
392 }
393 Ok(secret)
394}
395
396fn pure_ikm(ss_mlkem: &[u8], mlkem_ct: &[u8], key_hash: &[u8]) -> Vec<u8> {
397 encode(vec![bytes(ss_mlkem), bytes(mlkem_ct), bytes(key_hash)])
398}
399
400fn hybrid_ikm(
401 ss_mlkem: &[u8],
402 ss_ecdh: &[u8],
403 mlkem_ct: &[u8],
404 eph_pub: &[u8],
405 key_hash: &[u8],
406) -> Vec<u8> {
407 encode(vec![
408 bytes(ss_mlkem),
409 bytes(ss_ecdh),
410 bytes(mlkem_ct),
411 bytes(eph_pub),
412 bytes(key_hash),
413 ])
414}
415
416fn combine(label: &str, ikm: &[u8]) -> [u8; KEY_HASH_SIZE] {
418 extract(label.as_bytes(), ikm)
419}
420
421#[derive(Debug, Clone, Copy, PartialEq, Eq)]
424pub struct Parties {
425 pub request_id: [u8; 16],
426 pub caller: [u8; 32],
427 pub target: [u8; 32],
428}
429
430pub fn call_keys(
433 secret: &[u8; KEY_HASH_SIZE],
434 frame_type: &str,
435 p: &Parties,
436) -> ([u8; 32], [u8; 32]) {
437 halves(expand(
438 secret,
439 &encode(vec![
440 Value::text(LABEL_CALL),
441 Value::text(frame_type),
442 bytes(&p.request_id),
443 bytes(&p.caller),
444 bytes(&p.target),
445 ]),
446 ))
447}
448
449pub fn stream_keys(secret: &[u8; KEY_HASH_SIZE], p: &Parties) -> ([u8; 32], [u8; 32]) {
451 halves(expand(
452 secret,
453 &encode(vec![
454 Value::text(LABEL_STREAM),
455 bytes(&p.request_id),
456 bytes(&p.caller),
457 bytes(&p.target),
458 ]),
459 ))
460}
461
462fn halves(okm: [u8; 64]) -> ([u8; 32], [u8; 32]) {
463 let mut a = [0; 32];
464 let mut b = [0; 32];
465 a.copy_from_slice(&okm[..32]);
466 b.copy_from_slice(&okm[32..]);
467 (a, b)
468}
469
470#[derive(Debug, Clone, PartialEq, Eq)]
472pub struct Request {
473 pub frame_type: String,
474 pub realm: [u8; 32],
475 pub procedure: String,
476 pub caller: [u8; 32],
477 pub target: [u8; 32],
478 pub request_id: [u8; 16],
479 pub deadline: u64,
480}
481
482impl Request {
483 fn fields(&self, frame_type: &str) -> Vec<Value> {
484 vec![
485 Value::text(LABEL_AAD),
486 Value::text(frame_type),
487 bytes(&self.realm),
488 Value::text(self.procedure.clone()),
489 bytes(&self.caller),
490 bytes(&self.target),
491 bytes(&self.request_id),
492 Value::Int(i128::from(self.deadline)),
493 ]
494 }
495}
496
497pub fn request_aad(r: &Request) -> Vec<u8> {
500 encode(r.fields(&r.frame_type))
501}
502
503pub fn reply_aad(
507 r: &Request,
508 reply_frame_type: &str,
509 request_hash: &[u8; 48],
510 responded_by: &[u8; 32],
511) -> Vec<u8> {
512 let mut fields = r.fields(reply_frame_type);
513 fields.push(bytes(request_hash));
514 fields.push(bytes(responded_by));
515 encode(fields)
516}
517
518#[derive(Debug, Clone, Copy, PartialEq, Eq)]
520pub enum Direction {
521 CallerToProvider = 0,
523 ProviderToCaller = 1,
525}
526
527pub fn stream_aad(frame_type: &str, request_id: &[u8; 16], seq: u64, d: Direction) -> Vec<u8> {
529 encode(vec![
530 Value::text(LABEL_STREAM_AAD),
531 Value::text(frame_type),
532 bytes(request_id),
533 Value::Int(i128::from(seq)),
534 Value::Int(d as i128),
535 ])
536}
537
538pub fn stream_nonce(seq: u64) -> [u8; NONCE_SIZE] {
540 let mut nonce = [0; NONCE_SIZE];
541 nonce[4..].copy_from_slice(&seq.to_be_bytes());
542 nonce
543}
544
545pub fn random_nonce() -> Result<[u8; NONCE_SIZE], SealError> {
548 let mut nonce = [0; NONCE_SIZE];
549 aws_lc_rs::rand::fill(&mut nonce).map_err(|_| SealError::Unavailable)?;
550 Ok(nonce)
551}
552
553pub fn seal(key: &[u8; 32], nonce: &[u8; NONCE_SIZE], aad: &[u8], plain: &[u8]) -> Vec<u8> {
555 let key = LessSafeKey::new(UnboundKey::new(&AES_256_GCM, key).expect("a 32-byte AES-256 key"));
556 let mut out = plain.to_vec();
557 key.seal_in_place_append_tag(
558 Nonce::assume_unique_for_key(*nonce),
559 Aad::from(aad),
560 &mut out,
561 )
562 .expect("AES-256-GCM seals any payload under 2^36 bytes");
563 out
564}
565
566pub fn open(
569 key: &[u8; 32],
570 nonce: &[u8; NONCE_SIZE],
571 aad: &[u8],
572 sealed: &[u8],
573) -> Result<Vec<u8>, SealError> {
574 if sealed.len() < TAG_SIZE {
575 return Err(SealError::Refused);
576 }
577 let key = LessSafeKey::new(UnboundKey::new(&AES_256_GCM, key).expect("a 32-byte AES-256 key"));
578 let mut buf = sealed.to_vec();
579 let plain = key
580 .open_in_place(
581 Nonce::assume_unique_for_key(*nonce),
582 Aad::from(aad),
583 &mut buf,
584 )
585 .map_err(|_| SealError::Refused)?;
586 Ok(plain.to_vec())
587}
588
589pub fn error_plain(code: &str, detail: &str) -> Vec<u8> {
593 encode(vec![Value::text(code), Value::text(detail)])
594}
595
596pub fn open_error_plain(plain: &[u8]) -> Result<(String, Option<String>), SealError> {
599 match cbor::decode(plain) {
600 Ok(Value::List(items)) => match items.as_slice() {
601 [Value::Text(code), Value::Text(detail)]
602 if code.len() <= MAX_ERROR_CODE_BYTES && detail.len() <= MAX_ERROR_DETAIL_BYTES =>
603 {
604 Ok((code.clone(), (!detail.is_empty()).then(|| detail.clone())))
605 }
606 _ => Err(SealError::NotAnErrorPlain),
607 },
608 _ => Err(SealError::NotAnErrorPlain),
609 }
610}
611
612fn extract(salt: &[u8], ikm: &[u8]) -> [u8; KEY_HASH_SIZE] {
616 let tag = hmac::sign(&hmac::Key::new(hmac::HMAC_SHA384, salt), ikm);
617 let mut prk = [0; KEY_HASH_SIZE];
618 prk.copy_from_slice(tag.as_ref());
619 prk
620}
621
622fn expand(prk: &[u8; KEY_HASH_SIZE], info: &[u8]) -> [u8; 64] {
625 let key = hmac::Key::new(hmac::HMAC_SHA384, prk);
626 let mut okm = [0; 64];
627 let mut previous: Vec<u8> = Vec::new();
628 for (counter, chunk) in (1u8..).zip(okm.chunks_mut(KEY_HASH_SIZE)) {
629 let mut ctx = hmac::Context::with_key(&key);
630 ctx.update(&previous);
631 ctx.update(info);
632 ctx.update(&[counter]);
633 previous = ctx.sign().as_ref().to_vec();
634 chunk.copy_from_slice(&previous[..chunk.len()]);
635 }
636 okm
637}
638
639fn bytes(b: &[u8]) -> Value {
640 Value::Bytes(b.to_vec())
641}
642
643fn encode(items: Vec<Value>) -> Vec<u8> {
647 cbor::encode(&Value::List(items)).expect("seal arrays hold no integer outside u64")
648}
649
650fn hex(b: &[u8]) -> String {
651 b.iter().map(|x| format!("{x:02x}")).collect()
652}
653
654#[cfg(test)]
655mod vectors;