1use alloc::collections::{BTreeMap, VecDeque};
27use alloc::string::String;
28use alloc::vec::Vec;
29use core::fmt::Write as _;
30use core::mem;
31
32use chacha20poly1305::aead::{Aead, Payload};
33use chacha20poly1305::{Key, KeyInit, XChaCha20Poly1305, XNonce};
34use ed25519_dalek::{Signature, Signer, SigningKey, VerifyingKey};
35use hkdf::Hkdf;
36use hkdf::hmac::{Hmac, Mac};
37use sha2::{Digest, Sha256};
38use x25519_dalek::{PublicKey, StaticSecret};
39use zeroize::Zeroizing;
40
41pub const PROTOCOL_VERSION: u32 = 1;
44
45pub const MAX_SKIP: u32 = 1000;
48
49pub const MAX_SKIPPED_KEYS: usize = 2000;
53
54const PAD_BLOCK: usize = 64;
57
58const X3DH_INFO: &[u8] = b"obby.world/e2ee x3dh";
59const ROOT_INFO: &[u8] = b"obby.world/e2ee root";
60const MESSAGE_KEY_INFO: &[u8] = b"obby.world/e2ee message";
61const NONCE_INFO: &[u8] = b"obby.world/e2ee nonce";
62
63#[derive(Debug, Clone, Copy, PartialEq, Eq)]
66#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
67pub enum Error {
68 InvalidSignature,
70 NonContributoryDh,
73 Aead,
76 Padding,
78 TooManySkipped,
80 CounterOverflow,
83 NoChain,
85 FingerprintChanged {
88 previous: Fingerprint,
90 current: Fingerprint,
92 },
93 WrongState,
96 Fragmentation,
99 Internal,
102}
103
104pub trait RandomSource {
110 fn fill_bytes(&mut self, dest: &mut [u8]);
112}
113
114fn random_array<const N: usize>(rng: &mut impl RandomSource) -> [u8; N] {
115 let mut bytes = [0u8; N];
116 rng.fill_bytes(&mut bytes);
117 bytes
118}
119
120fn generate_x25519_keypair(rng: &mut impl RandomSource) -> ([u8; 32], [u8; 32]) {
121 let secret = random_array::<32>(rng);
122 let public = *PublicKey::from(&StaticSecret::from(secret)).as_bytes();
123 (secret, public)
124}
125
126fn concat(parts: &[&[u8]]) -> Vec<u8> {
127 let mut out = Vec::new();
128 for part in parts {
129 out.extend_from_slice(part);
130 }
131 out
132}
133
134fn hkdf_sha256(salt: &[u8], ikm: &[u8], info: &[u8], out: &mut [u8]) -> Result<(), Error> {
135 let hk = Hkdf::<Sha256>::new(Some(salt), ikm);
136 hk.expand(info, out).map_err(|_| Error::Internal)
137}
138
139fn hmac_sha256(key: &[u8; 32], data: &[u8]) -> Result<[u8; 32], Error> {
140 let mut mac = <Hmac<Sha256> as Mac>::new_from_slice(key).map_err(|_| Error::Internal)?;
141 mac.update(data);
142 Ok(mac.finalize().into_bytes().into())
143}
144
145fn diffie_hellman_raw(secret: &[u8; 32], public: &[u8; 32]) -> Result<[u8; 32], Error> {
146 let secret = StaticSecret::from(*secret);
147 let public = PublicKey::from(*public);
148 let shared = secret.diffie_hellman(&public);
149 if !shared.was_contributory() {
150 return Err(Error::NonContributoryDh);
151 }
152 Ok(*shared.as_bytes())
153}
154
155#[derive(Debug, Clone, Copy, PartialEq, Eq)]
162#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
163pub struct IdentityPublic {
164 pub agreement: [u8; 32],
166 pub signing: [u8; 32],
169}
170
171struct IdentitySecret {
172 agreement: Zeroizing<[u8; 32]>,
173 signing: Zeroizing<[u8; 32]>,
174}
175
176pub struct Identity {
179 secret: IdentitySecret,
180 public: IdentityPublic,
181}
182
183impl Identity {
184 pub fn generate(rng: &mut impl RandomSource) -> Self {
186 let agreement_secret = random_array::<32>(rng);
187 let signing_secret = random_array::<32>(rng);
188 let agreement_public = *PublicKey::from(&StaticSecret::from(agreement_secret)).as_bytes();
189 let signing_public = *SigningKey::from_bytes(&signing_secret)
190 .verifying_key()
191 .as_bytes();
192 Self {
193 secret: IdentitySecret {
194 agreement: Zeroizing::new(agreement_secret),
195 signing: Zeroizing::new(signing_secret),
196 },
197 public: IdentityPublic {
198 agreement: agreement_public,
199 signing: signing_public,
200 },
201 }
202 }
203
204 pub const fn public(&self) -> IdentityPublic {
206 self.public
207 }
208
209 pub fn fingerprint(&self) -> Fingerprint {
211 Fingerprint::of_signing_key(&self.public.signing)
212 }
213
214 fn sign(&self, message: &[u8]) -> Result<Vec<u8>, Error> {
215 let signing_key = SigningKey::from_bytes(&self.secret.signing);
216 let signature: Signature = signing_key
217 .try_sign(message)
218 .map_err(|_| Error::InvalidSignature)?;
219 Ok(signature.to_bytes().to_vec())
220 }
221}
222
223#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
228#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
229pub struct Fingerprint([u8; 16]);
230
231impl Fingerprint {
232 pub fn of_signing_key(signing_public: &[u8; 32]) -> Self {
234 let digest = Sha256::digest(signing_public);
235 let (head, _tail) = digest.split_at(16);
236 let mut bytes = [0u8; 16];
237 bytes.copy_from_slice(head);
238 Self(bytes)
239 }
240
241 pub fn safety_number(&self) -> String {
244 let mut out = String::with_capacity(39);
245 for (index, pair) in self.0.chunks(2).enumerate() {
246 if index > 0 {
247 out.push(' ');
248 }
249 for byte in pair {
250 let _ = write!(out, "{byte:02X}");
251 }
252 }
253 out
254 }
255}
256
257#[derive(Debug, Clone, Copy, PartialEq, Eq)]
259pub enum PinOutcome {
260 New,
262 Same,
264 Changed {
267 previous: Fingerprint,
269 },
270}
271
272#[derive(Debug, Clone, Default)]
277pub struct PeerTrust {
278 pinned: Option<Fingerprint>,
279 verified: bool,
280}
281
282impl PeerTrust {
283 pub fn new() -> Self {
285 Self::default()
286 }
287
288 pub fn observe(&mut self, fingerprint: Fingerprint) -> PinOutcome {
290 match self.pinned {
291 None => {
292 self.pinned = Some(fingerprint);
293 PinOutcome::New
294 }
295 Some(pinned) if pinned == fingerprint => PinOutcome::Same,
296 Some(previous) => PinOutcome::Changed { previous },
297 }
298 }
299
300 pub fn repin(&mut self, fingerprint: Fingerprint) {
303 self.pinned = Some(fingerprint);
304 self.verified = false;
305 }
306
307 pub const fn pinned(&self) -> Option<Fingerprint> {
309 self.pinned
310 }
311
312 pub const fn is_verified(&self) -> bool {
314 self.verified
315 }
316
317 pub fn set_verified(&mut self, verified: bool) {
319 self.verified = verified;
320 }
321}
322
323pub fn keeps_own_offer(own: Fingerprint, peer: Option<Fingerprint>) -> bool {
330 match peer {
331 Some(peer) => own < peer,
332 None => false,
333 }
334}
335
336#[derive(Debug, Clone, PartialEq, Eq)]
343#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
344pub struct PreKeyBundle {
345 pub ik: [u8; 32],
347 pub sik: [u8; 32],
349 pub spk: [u8; 32],
351 pub sig: Vec<u8>,
353 pub opk: [u8; 32],
355}
356
357#[derive(Debug, Clone, PartialEq, Eq)]
360#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
361pub struct HandshakeResponse {
362 pub ik: [u8; 32],
364 pub sik: [u8; 32],
366 pub ek: [u8; 32],
368 pub sig: Vec<u8>,
370 pub boot: RatchetMessage,
373}
374
375#[derive(Debug, Clone, PartialEq, Eq)]
378#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
379pub struct RatchetMessage {
380 pub dh: [u8; 32],
382 pub pn: u32,
384 pub n: u32,
386 pub ct: Vec<u8>,
388}
389
390#[derive(Debug, Clone, PartialEq, Eq)]
396#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
397pub enum Frame {
398 Init {
400 bundle: PreKeyBundle,
402 account: Option<String>,
404 },
405 Accept {
407 response: HandshakeResponse,
409 account: Option<String>,
411 },
412 Reject {
414 reason: Option<String>,
416 },
417 Ack {
419 ct: RatchetMessage,
421 },
422 Close,
424 Msg {
426 ct: RatchetMessage,
428 },
429 Media {
431 ct: RatchetMessage,
433 },
434}
435
436#[derive(Debug, Clone, PartialEq, Eq)]
439#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
440pub struct Frag {
441 pub id: String,
443 pub i: u32,
445 pub n: u32,
447 pub ct: Vec<u8>,
449}
450
451pub fn reassemble(fragments: &[Frag]) -> Result<Vec<u8>, Error> {
458 let first = fragments.first().ok_or(Error::Fragmentation)?;
459 let id = &first.id;
460 let total = first.n;
461 let expected = u32::try_from(fragments.len()).map_err(|_| Error::Fragmentation)?;
462 if expected != total {
463 return Err(Error::Fragmentation);
464 }
465
466 let mut slots: Vec<Option<&[u8]>> = alloc::vec![None; fragments.len()];
467 for fragment in fragments {
468 if &fragment.id != id || fragment.n != total {
469 return Err(Error::Fragmentation);
470 }
471 let index = usize::try_from(fragment.i).map_err(|_| Error::Fragmentation)?;
472 let slot = slots.get_mut(index).ok_or(Error::Fragmentation)?;
473 if slot.is_some() {
474 return Err(Error::Fragmentation);
475 }
476 *slot = Some(&fragment.ct);
477 }
478
479 let mut out = Vec::new();
480 for slot in slots {
481 out.extend_from_slice(slot.ok_or(Error::Fragmentation)?);
482 }
483 Ok(out)
484}
485
486pub struct PendingOffer {
493 bundle: PreKeyBundle,
494 spk_secret: Zeroizing<[u8; 32]>,
495 opk_secret: Zeroizing<[u8; 32]>,
496}
497
498pub fn create_offer(
501 identity: &Identity,
502 rng: &mut impl RandomSource,
503) -> Result<PendingOffer, Error> {
504 let spk_secret = random_array::<32>(rng);
505 let opk_secret = random_array::<32>(rng);
506 let spk_public = *PublicKey::from(&StaticSecret::from(spk_secret)).as_bytes();
507 let opk_public = *PublicKey::from(&StaticSecret::from(opk_secret)).as_bytes();
508
509 let public = identity.public();
510 let signed = concat(&[&public.agreement, &spk_public, &opk_public]);
511 let sig = identity.sign(&signed)?;
512
513 Ok(PendingOffer {
514 bundle: PreKeyBundle {
515 ik: public.agreement,
516 sik: public.signing,
517 spk: spk_public,
518 sig,
519 opk: opk_public,
520 },
521 spk_secret: Zeroizing::new(spk_secret),
522 opk_secret: Zeroizing::new(opk_secret),
523 })
524}
525
526fn verify_signature(
527 signing_public: &[u8; 32],
528 message: &[u8],
529 signature: &[u8],
530) -> Result<(), Error> {
531 let verifying_key =
532 VerifyingKey::from_bytes(signing_public).map_err(|_| Error::InvalidSignature)?;
533 let signature = Signature::try_from(signature).map_err(|_| Error::InvalidSignature)?;
534 verifying_key
535 .verify_strict(message, &signature)
536 .map_err(|_| Error::InvalidSignature)
537}
538
539fn verify_bundle_signature(bundle: &PreKeyBundle) -> Result<(), Error> {
540 let signed = concat(&[&bundle.ik, &bundle.spk, &bundle.opk]);
541 verify_signature(&bundle.sik, &signed, &bundle.sig)
542}
543
544fn x3dh_kdf(
545 dh1: &[u8; 32],
546 dh2: &[u8; 32],
547 dh3: &[u8; 32],
548 dh4: &[u8; 32],
549) -> Result<Zeroizing<[u8; 32]>, Error> {
550 let ikm = concat(&[dh1, dh2, dh3, dh4]);
551 let mut sk = [0u8; 32];
552 hkdf_sha256(&[0u8; 32], &ikm, X3DH_INFO, &mut sk)?;
553 Ok(Zeroizing::new(sk))
554}
555
556fn x3dh_secret_responder(
558 own_identity: &[u8; 32],
559 own_ephemeral: &[u8; 32],
560 peer_identity: &[u8; 32],
561 peer_signed_prekey: &[u8; 32],
562 peer_one_time_prekey: &[u8; 32],
563) -> Result<Zeroizing<[u8; 32]>, Error> {
564 let dh1 = diffie_hellman_raw(own_identity, peer_signed_prekey)?;
565 let dh2 = diffie_hellman_raw(own_ephemeral, peer_identity)?;
566 let dh3 = diffie_hellman_raw(own_ephemeral, peer_signed_prekey)?;
567 let dh4 = diffie_hellman_raw(own_ephemeral, peer_one_time_prekey)?;
568 x3dh_kdf(&dh1, &dh2, &dh3, &dh4)
569}
570
571fn x3dh_secret_initiator(
574 own_signed_prekey: &[u8; 32],
575 own_identity: &[u8; 32],
576 own_one_time_prekey: &[u8; 32],
577 peer_identity: &[u8; 32],
578 peer_ephemeral: &[u8; 32],
579) -> Result<Zeroizing<[u8; 32]>, Error> {
580 let dh1 = diffie_hellman_raw(own_signed_prekey, peer_identity)?;
581 let dh2 = diffie_hellman_raw(own_identity, peer_ephemeral)?;
582 let dh3 = diffie_hellman_raw(own_signed_prekey, peer_ephemeral)?;
583 let dh4 = diffie_hellman_raw(own_one_time_prekey, peer_ephemeral)?;
584 x3dh_kdf(&dh1, &dh2, &dh3, &dh4)
585}
586
587pub fn accept_offer(
594 identity: &Identity,
595 bundle: &PreKeyBundle,
596 rng: &mut impl RandomSource,
597) -> Result<(HandshakeResponse, Ratchet), Error> {
598 verify_bundle_signature(bundle)?;
599
600 let (ek_secret, ek_public) = generate_x25519_keypair(rng);
601 let sk = x3dh_secret_responder(
602 &identity.secret.agreement,
603 &ek_secret,
604 &bundle.ik,
605 &bundle.spk,
606 &bundle.opk,
607 )?;
608
609 let mut ratchet = Ratchet::init_as_responder(*sk, ek_secret, ek_public, bundle.spk)?;
610 let boot = ratchet.encrypt(&[])?;
611
612 let public = identity.public();
613 let signed = concat(&[&public.agreement, &ek_public]);
614 let sig = identity.sign(&signed)?;
615
616 let response = HandshakeResponse {
617 ik: public.agreement,
618 sik: public.signing,
619 ek: ek_public,
620 sig,
621 boot,
622 };
623 Ok((response, ratchet))
624}
625
626pub fn complete_handshake(
635 identity: &Identity,
636 pending: &PendingOffer,
637 response: &HandshakeResponse,
638 rng: &mut impl RandomSource,
639) -> Result<Ratchet, Error> {
640 let signed = concat(&[&response.ik, &response.ek]);
641 verify_signature(&response.sik, &signed, &response.sig)?;
642
643 let sk = x3dh_secret_initiator(
644 &pending.spk_secret,
645 &identity.secret.agreement,
646 &pending.opk_secret,
647 &response.ik,
648 &response.ek,
649 )?;
650
651 let mut ratchet = Ratchet::init_as_initiator(*sk, *pending.spk_secret, pending.bundle.spk);
652 ratchet.decrypt(&response.boot, rng)?;
653 Ok(ratchet)
654}
655
656#[derive(Clone)]
661struct SkippedKeys {
662 by_id: BTreeMap<([u8; 32], u32), Zeroizing<[u8; 32]>>,
663 order: VecDeque<([u8; 32], u32)>,
664}
665
666impl SkippedKeys {
667 fn new() -> Self {
668 Self {
669 by_id: BTreeMap::new(),
670 order: VecDeque::new(),
671 }
672 }
673
674 fn insert(&mut self, dh: [u8; 32], n: u32, key: Zeroizing<[u8; 32]>) {
675 let id = (dh, n);
676 if self.by_id.insert(id, key).is_none() {
677 self.order.push_back(id);
678 }
679 while self.order.len() > MAX_SKIPPED_KEYS {
680 if let Some(oldest) = self.order.pop_front() {
681 self.by_id.remove(&oldest);
682 }
683 }
684 }
685
686 fn take(&mut self, dh: [u8; 32], n: u32) -> Option<Zeroizing<[u8; 32]>> {
687 let id = (dh, n);
688 let key = self.by_id.remove(&id)?;
689 self.order.retain(|entry| *entry != id);
690 Some(key)
691 }
692}
693
694#[derive(Clone)]
697pub struct Ratchet {
698 root_key: Zeroizing<[u8; 32]>,
699 dhs_secret: Zeroizing<[u8; 32]>,
700 dhs_public: [u8; 32],
701 dhr: Option<[u8; 32]>,
702 send_chain: Option<Zeroizing<[u8; 32]>>,
703 recv_chain: Option<Zeroizing<[u8; 32]>>,
704 n_send: u32,
705 n_recv: u32,
706 prev_chain_len: u32,
707 skipped: SkippedKeys,
708}
709
710impl Ratchet {
711 fn init_as_responder(
714 root_key: [u8; 32],
715 own_ek_secret: [u8; 32],
716 own_ek_public: [u8; 32],
717 their_spk_public: [u8; 32],
718 ) -> Result<Self, Error> {
719 let dh_out = diffie_hellman_raw(&own_ek_secret, &their_spk_public)?;
720 let (new_root, send_chain) = kdf_root(&root_key, &dh_out)?;
721 Ok(Self {
722 root_key: Zeroizing::new(new_root),
723 dhs_secret: Zeroizing::new(own_ek_secret),
724 dhs_public: own_ek_public,
725 dhr: Some(their_spk_public),
726 send_chain: Some(Zeroizing::new(send_chain)),
727 recv_chain: None,
728 n_send: 0,
729 n_recv: 0,
730 prev_chain_len: 0,
731 skipped: SkippedKeys::new(),
732 })
733 }
734
735 fn init_as_initiator(
739 root_key: [u8; 32],
740 own_dhs_secret: [u8; 32],
741 own_dhs_public: [u8; 32],
742 ) -> Self {
743 Self {
744 root_key: Zeroizing::new(root_key),
745 dhs_secret: Zeroizing::new(own_dhs_secret),
746 dhs_public: own_dhs_public,
747 dhr: None,
748 send_chain: None,
749 recv_chain: None,
750 n_send: 0,
751 n_recv: 0,
752 prev_chain_len: 0,
753 skipped: SkippedKeys::new(),
754 }
755 }
756
757 pub fn encrypt(&mut self, plaintext: &[u8]) -> Result<RatchetMessage, Error> {
759 let Some(chain) = self.send_chain.clone() else {
760 return Err(Error::NoChain);
761 };
762 let (message_key, next_chain) = kdf_chain(&chain)?;
763 self.send_chain = Some(Zeroizing::new(next_chain));
764
765 let dh = self.dhs_public;
766 let pn = self.prev_chain_len;
767 let n = self.n_send;
768 self.n_send = self.n_send.checked_add(1).ok_or(Error::CounterOverflow)?;
769
770 let padded = pad(plaintext);
771 let aad = header_aad(&dh, pn, n);
772 let ct = aead_encrypt(&message_key, &aad, &padded)?;
773 Ok(RatchetMessage { dh, pn, n, ct })
774 }
775
776 pub fn decrypt(
783 &mut self,
784 msg: &RatchetMessage,
785 rng: &mut impl RandomSource,
786 ) -> Result<Vec<u8>, Error> {
787 let mut trial = self.clone();
788
789 if let Some(message_key) = trial.skipped.take(msg.dh, msg.n) {
790 let plaintext = decrypt_with_key(&message_key, msg)?;
791 *self = trial;
792 return Ok(plaintext);
793 }
794
795 if trial.dhr != Some(msg.dh) {
796 trial.skip_current_receiving_chain(msg.pn)?;
797 trial.dh_ratchet_receive(msg.dh, rng)?;
798 }
799 trial.skip_current_receiving_chain(msg.n)?;
800
801 let Some(chain) = trial.recv_chain.clone() else {
802 return Err(Error::NoChain);
803 };
804 let (message_key, next_chain) = kdf_chain(&chain)?;
805 trial.recv_chain = Some(Zeroizing::new(next_chain));
806 trial.n_recv = trial.n_recv.checked_add(1).ok_or(Error::CounterOverflow)?;
807
808 let plaintext = decrypt_with_key(&Zeroizing::new(message_key), msg)?;
809 *self = trial;
810 Ok(plaintext)
811 }
812
813 fn skip_current_receiving_chain(&mut self, until: u32) -> Result<(), Error> {
817 let Some(dhr) = self.dhr else { return Ok(()) };
818 let Some(mut chain) = self.recv_chain.take() else {
819 return Ok(());
820 };
821 if until <= self.n_recv {
822 self.recv_chain = Some(chain);
823 return Ok(());
824 }
825 let span = until - self.n_recv;
826 if span > MAX_SKIP {
827 self.recv_chain = Some(chain);
828 return Err(Error::TooManySkipped);
829 }
830 for _ in 0..span {
831 let (message_key, next_chain) = kdf_chain(&chain)?;
832 self.skipped
833 .insert(dhr, self.n_recv, Zeroizing::new(message_key));
834 chain = Zeroizing::new(next_chain);
835 self.n_recv = self.n_recv.checked_add(1).ok_or(Error::CounterOverflow)?;
836 }
837 self.recv_chain = Some(chain);
838 Ok(())
839 }
840
841 fn dh_ratchet_receive(
845 &mut self,
846 new_dhr: [u8; 32],
847 rng: &mut impl RandomSource,
848 ) -> Result<(), Error> {
849 self.prev_chain_len = self.n_send;
850 self.n_send = 0;
851 self.n_recv = 0;
852 self.dhr = Some(new_dhr);
853
854 let dh_out = diffie_hellman_raw(&self.dhs_secret, &new_dhr)?;
855 let (root_after_recv, recv_chain) = kdf_root(&self.root_key, &dh_out)?;
856 self.root_key = Zeroizing::new(root_after_recv);
857 self.recv_chain = Some(Zeroizing::new(recv_chain));
858
859 let (dhs_secret, dhs_public) = generate_x25519_keypair(rng);
860 self.dhs_secret = Zeroizing::new(dhs_secret);
861 self.dhs_public = dhs_public;
862
863 let dh_out2 = diffie_hellman_raw(&self.dhs_secret, &new_dhr)?;
864 let (root_after_send, send_chain) = kdf_root(&self.root_key, &dh_out2)?;
865 self.root_key = Zeroizing::new(root_after_send);
866 self.send_chain = Some(Zeroizing::new(send_chain));
867 Ok(())
868 }
869}
870
871fn decrypt_with_key(message_key: &[u8; 32], msg: &RatchetMessage) -> Result<Vec<u8>, Error> {
872 let aad = header_aad(&msg.dh, msg.pn, msg.n);
873 let padded = aead_decrypt(message_key, &aad, &msg.ct)?;
874 unpad(&padded)
875}
876
877fn kdf_root(root_key: &[u8; 32], dh_out: &[u8; 32]) -> Result<([u8; 32], [u8; 32]), Error> {
878 let mut okm = [0u8; 64];
879 hkdf_sha256(root_key, dh_out, ROOT_INFO, &mut okm)?;
880 let (root_half, chain_half) = okm.split_at(32);
881 let new_root: [u8; 32] = root_half.try_into().map_err(|_| Error::Internal)?;
882 let new_chain: [u8; 32] = chain_half.try_into().map_err(|_| Error::Internal)?;
883 Ok((new_root, new_chain))
884}
885
886fn kdf_chain(chain_key: &[u8; 32]) -> Result<([u8; 32], [u8; 32]), Error> {
887 let message_key = hmac_sha256(chain_key, &[0x01])?;
888 let next_chain = hmac_sha256(chain_key, &[0x02])?;
889 Ok((message_key, next_chain))
890}
891
892fn message_aead_params(message_key: &[u8; 32]) -> Result<(Key, XNonce), Error> {
893 let mut key_bytes = [0u8; 32];
894 hkdf_sha256(&[0u8; 32], message_key, MESSAGE_KEY_INFO, &mut key_bytes)?;
895 let mut nonce_bytes = [0u8; 24];
896 hkdf_sha256(&[0u8; 32], message_key, NONCE_INFO, &mut nonce_bytes)?;
897 Ok((Key::from(key_bytes), XNonce::from(nonce_bytes)))
898}
899
900fn aead_encrypt(message_key: &[u8; 32], aad: &[u8], plaintext: &[u8]) -> Result<Vec<u8>, Error> {
901 let (key, nonce) = message_aead_params(message_key)?;
902 XChaCha20Poly1305::new(&key)
903 .encrypt(
904 &nonce,
905 Payload {
906 msg: plaintext,
907 aad,
908 },
909 )
910 .map_err(|_| Error::Aead)
911}
912
913fn aead_decrypt(message_key: &[u8; 32], aad: &[u8], ciphertext: &[u8]) -> Result<Vec<u8>, Error> {
914 let (key, nonce) = message_aead_params(message_key)?;
915 XChaCha20Poly1305::new(&key)
916 .decrypt(
917 &nonce,
918 Payload {
919 msg: ciphertext,
920 aad,
921 },
922 )
923 .map_err(|_| Error::Aead)
924}
925
926fn header_aad(dh: &[u8; 32], pn: u32, n: u32) -> Vec<u8> {
927 let mut aad = Vec::with_capacity(40);
928 aad.extend_from_slice(dh);
929 aad.extend_from_slice(&pn.to_be_bytes());
930 aad.extend_from_slice(&n.to_be_bytes());
931 aad
932}
933
934fn pad(plaintext: &[u8]) -> Vec<u8> {
935 let mut out = Vec::with_capacity(plaintext.len() + PAD_BLOCK);
936 out.extend_from_slice(plaintext);
937 out.push(0x80);
938 let remainder = out.len() % PAD_BLOCK;
939 if remainder != 0 {
940 out.resize(out.len() + (PAD_BLOCK - remainder), 0);
941 }
942 out
943}
944
945fn unpad(padded: &[u8]) -> Result<Vec<u8>, Error> {
946 let marker = padded
947 .iter()
948 .rposition(|&byte| byte != 0)
949 .ok_or(Error::Padding)?;
950 if padded.get(marker) != Some(&0x80) {
951 return Err(Error::Padding);
952 }
953 padded
954 .get(..marker)
955 .map(<[u8]>::to_vec)
956 .ok_or(Error::Padding)
957}
958
959#[derive(Debug, Clone, Copy, PartialEq, Eq)]
965pub enum Role {
966 Initiator,
968 Responder,
970}
971
972pub enum SessionState {
980 Idle,
982 Offered {
984 pending: PendingOffer,
986 },
987 OfferReceived {
989 bundle: PreKeyBundle,
991 },
992 AwaitingAck {
997 ratchet: Ratchet,
999 },
1000 Established {
1002 ratchet: Ratchet,
1004 },
1005 Rejected,
1007 Closed,
1009}
1010
1011pub struct Session {
1019 state: SessionState,
1020 trust: PeerTrust,
1021 role: Option<Role>,
1022}
1023
1024impl Default for Session {
1025 fn default() -> Self {
1026 Self {
1027 state: SessionState::Idle,
1028 trust: PeerTrust::new(),
1029 role: None,
1030 }
1031 }
1032}
1033
1034impl Session {
1035 pub fn new() -> Self {
1037 Self::default()
1038 }
1039
1040 pub fn start(
1042 &mut self,
1043 identity: &Identity,
1044 rng: &mut impl RandomSource,
1045 ) -> Result<PreKeyBundle, Error> {
1046 let pending = create_offer(identity, rng)?;
1047 let bundle = pending.bundle.clone();
1048 self.state = SessionState::Offered { pending };
1049 Ok(bundle)
1050 }
1051
1052 pub fn receive_offer(&mut self, bundle: PreKeyBundle) -> Fingerprint {
1059 let offered = Fingerprint::of_signing_key(&bundle.sik);
1060 self.state = SessionState::OfferReceived { bundle };
1061 offered
1062 }
1063
1064 pub fn accept(
1067 &mut self,
1068 identity: &Identity,
1069 rng: &mut impl RandomSource,
1070 ) -> Result<HandshakeResponse, Error> {
1071 let SessionState::OfferReceived { bundle } = &self.state else {
1072 return Err(Error::WrongState);
1073 };
1074 let (response, ratchet) = accept_offer(identity, bundle, rng)?;
1075
1076 let fingerprint = Fingerprint::of_signing_key(&bundle.sik);
1077 if let PinOutcome::Changed { previous } = self.trust.observe(fingerprint) {
1078 return Err(Error::FingerprintChanged {
1079 previous,
1080 current: fingerprint,
1081 });
1082 }
1083
1084 self.role = Some(Role::Responder);
1085 self.state = SessionState::AwaitingAck { ratchet };
1086 Ok(response)
1087 }
1088
1089 pub fn reject(&mut self, reason: Option<String>) -> Frame {
1091 self.state = SessionState::Rejected;
1092 Frame::Reject { reason }
1093 }
1094
1095 pub fn receive_reject(&mut self) {
1097 self.state = SessionState::Rejected;
1098 }
1099
1100 pub fn receive_accept(
1105 &mut self,
1106 identity: &Identity,
1107 response: &HandshakeResponse,
1108 rng: &mut impl RandomSource,
1109 ) -> Result<(), Error> {
1110 let SessionState::Offered { pending } = &self.state else {
1111 return Err(Error::WrongState);
1112 };
1113 let ratchet = complete_handshake(identity, pending, response, rng)?;
1114
1115 let fingerprint = Fingerprint::of_signing_key(&response.sik);
1116 if let PinOutcome::Changed { previous } = self.trust.observe(fingerprint) {
1117 return Err(Error::FingerprintChanged {
1118 previous,
1119 current: fingerprint,
1120 });
1121 }
1122
1123 self.role = Some(Role::Initiator);
1124 self.state = SessionState::Established { ratchet };
1125 Ok(())
1126 }
1127
1128 pub fn make_ack(&mut self) -> Result<RatchetMessage, Error> {
1131 self.send(&[])
1132 }
1133
1134 pub fn receive_ack(
1138 &mut self,
1139 ct: &RatchetMessage,
1140 rng: &mut impl RandomSource,
1141 ) -> Result<(), Error> {
1142 let SessionState::AwaitingAck { ratchet } = &mut self.state else {
1143 return Err(Error::WrongState);
1144 };
1145 ratchet.decrypt(ct, rng)?;
1146 let SessionState::AwaitingAck { ratchet } =
1147 mem::replace(&mut self.state, SessionState::Idle)
1148 else {
1149 return Err(Error::WrongState);
1150 };
1151 self.state = SessionState::Established { ratchet };
1152 Ok(())
1153 }
1154
1155 pub fn send(&mut self, plaintext: &[u8]) -> Result<RatchetMessage, Error> {
1157 let SessionState::Established { ratchet } = &mut self.state else {
1158 return Err(Error::WrongState);
1159 };
1160 ratchet.encrypt(plaintext)
1161 }
1162
1163 pub fn receive(
1166 &mut self,
1167 ct: &RatchetMessage,
1168 rng: &mut impl RandomSource,
1169 ) -> Result<Vec<u8>, Error> {
1170 let SessionState::Established { ratchet } = &mut self.state else {
1171 return Err(Error::WrongState);
1172 };
1173 ratchet.decrypt(ct, rng)
1174 }
1175
1176 pub fn close(&mut self) -> Frame {
1178 self.state = SessionState::Closed;
1179 Frame::Close
1180 }
1181
1182 pub fn receive_close(&mut self) {
1184 self.state = SessionState::Closed;
1185 }
1186
1187 pub const fn is_established(&self) -> bool {
1189 matches!(self.state, SessionState::Established { .. })
1190 }
1191
1192 pub fn peer_fingerprint(&self) -> Option<Fingerprint> {
1194 self.trust.pinned()
1195 }
1196
1197 pub fn is_peer_verified(&self) -> bool {
1199 self.trust.is_verified()
1200 }
1201
1202 pub fn mark_peer_verified(&mut self) {
1204 self.trust.set_verified(true);
1205 }
1206
1207 pub fn confirm_fingerprint_change(&mut self, fingerprint: Fingerprint) {
1210 self.trust.repin(fingerprint);
1211 }
1212
1213 pub const fn role(&self) -> Option<Role> {
1215 self.role
1216 }
1217}
1218
1219#[cfg(test)]
1220mod tests {
1221 use super::*;
1222
1223 struct TestRng(u64);
1224
1225 impl TestRng {
1226 fn seeded(seed: u64) -> Self {
1227 Self(seed)
1228 }
1229 }
1230
1231 impl RandomSource for TestRng {
1232 fn fill_bytes(&mut self, dest: &mut [u8]) {
1233 for chunk in dest.chunks_mut(8) {
1234 self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
1235 let mut z = self.0;
1236 z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
1237 z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
1238 z ^= z >> 31;
1239 let bytes = z.to_le_bytes();
1240 chunk.copy_from_slice(&bytes[..chunk.len()]);
1241 }
1242 }
1243 }
1244
1245 struct Pair {
1246 alice_identity: Identity,
1247 bob_identity: Identity,
1248 alice: Session,
1249 bob: Session,
1250 }
1251
1252 fn establish() -> Pair {
1253 let mut rng = TestRng::seeded(1);
1254 let alice_identity = Identity::generate(&mut rng);
1255 let bob_identity = Identity::generate(&mut rng);
1256 let mut alice = Session::new();
1257 let mut bob = Session::new();
1258
1259 let bundle = alice.start(&alice_identity, &mut rng).unwrap();
1260 bob.receive_offer(bundle);
1261 let response = bob.accept(&bob_identity, &mut rng).unwrap();
1262 assert!(!bob.is_established());
1263
1264 alice
1265 .receive_accept(&alice_identity, &response, &mut rng)
1266 .unwrap();
1267 assert!(alice.is_established());
1268
1269 let ack = alice.make_ack().unwrap();
1270 bob.receive_ack(&ack, &mut rng).unwrap();
1271 assert!(bob.is_established());
1272
1273 Pair {
1274 alice_identity,
1275 bob_identity,
1276 alice,
1277 bob,
1278 }
1279 }
1280
1281 #[test]
1282 fn full_handshake_establishes_both_sides() {
1283 let pair = establish();
1284 assert_eq!(pair.alice.role(), Some(Role::Initiator));
1285 assert_eq!(pair.bob.role(), Some(Role::Responder));
1286 assert_eq!(
1287 pair.alice.peer_fingerprint(),
1288 Some(pair.bob_identity.fingerprint())
1289 );
1290 assert_eq!(
1291 pair.bob.peer_fingerprint(),
1292 Some(pair.alice_identity.fingerprint())
1293 );
1294 }
1295
1296 #[test]
1297 fn accept_alone_does_not_establish() {
1298 let mut rng = TestRng::seeded(2);
1299 let alice_identity = Identity::generate(&mut rng);
1300 let bob_identity = Identity::generate(&mut rng);
1301 let mut alice = Session::new();
1302 let mut bob = Session::new();
1303
1304 let bundle = alice.start(&alice_identity, &mut rng).unwrap();
1305 bob.receive_offer(bundle);
1306 bob.accept(&bob_identity, &mut rng).unwrap();
1307
1308 assert!(!bob.is_established());
1309 assert_eq!(bob.send(b"hello"), Err(Error::WrongState));
1310 }
1311
1312 #[test]
1313 fn established_session_encrypts_and_decrypts() {
1314 let mut pair = establish();
1315 let mut rng = TestRng::seeded(3);
1316
1317 let ct = pair.alice.send(b"hello bob").unwrap();
1318 let plaintext = pair.bob.receive(&ct, &mut rng).unwrap();
1319 assert_eq!(plaintext, b"hello bob");
1320
1321 let ct = pair.bob.send(b"hello alice").unwrap();
1322 let plaintext = pair.alice.receive(&ct, &mut rng).unwrap();
1323 assert_eq!(plaintext, b"hello alice");
1324 }
1325
1326 #[test]
1327 fn out_of_order_delivery_still_decrypts() {
1328 let mut pair = establish();
1329 let mut rng = TestRng::seeded(4);
1330
1331 let first = pair.alice.send(b"one").unwrap();
1332 let second = pair.alice.send(b"two").unwrap();
1333 let third = pair.alice.send(b"three").unwrap();
1334
1335 assert_eq!(pair.bob.receive(&third, &mut rng).unwrap(), b"three");
1336 assert_eq!(pair.bob.receive(&first, &mut rng).unwrap(), b"one");
1337 assert_eq!(pair.bob.receive(&second, &mut rng).unwrap(), b"two");
1338 }
1339
1340 #[test]
1341 fn skip_bound_is_enforced() {
1342 let mut pair = establish();
1343 let mut rng = TestRng::seeded(5);
1344
1345 let mut far = pair.alice.send(b"far").unwrap();
1346 far.n += MAX_SKIP + 1;
1347 assert_eq!(pair.bob.receive(&far, &mut rng), Err(Error::TooManySkipped));
1348 }
1349
1350 #[test]
1351 fn tampered_ciphertext_fails_to_decrypt() {
1352 let mut pair = establish();
1353 let mut rng = TestRng::seeded(6);
1354
1355 let mut ct = pair.alice.send(b"hello").unwrap();
1356 let last = ct.ct.len() - 1;
1357 if let Some(byte) = ct.ct.get_mut(last) {
1358 *byte ^= 0xFF;
1359 }
1360 assert_eq!(pair.bob.receive(&ct, &mut rng), Err(Error::Aead));
1361 }
1362
1363 #[test]
1364 fn forged_responder_signature_is_rejected() {
1365 let mut rng = TestRng::seeded(7);
1366 let alice_identity = Identity::generate(&mut rng);
1367 let bob_identity = Identity::generate(&mut rng);
1368 let mallory_identity = Identity::generate(&mut rng);
1369 let mut alice = Session::new();
1370 let mut bob = Session::new();
1371
1372 let bundle = alice.start(&alice_identity, &mut rng).unwrap();
1373 bob.receive_offer(bundle);
1374 let mut response = bob.accept(&bob_identity, &mut rng).unwrap();
1375
1376 response.sik = mallory_identity.public().signing;
1379
1380 let result = alice.receive_accept(&alice_identity, &response, &mut rng);
1381 assert_eq!(result, Err(Error::InvalidSignature));
1382 assert!(!alice.is_established());
1383 assert_eq!(alice.peer_fingerprint(), None);
1384 }
1385
1386 #[test]
1387 fn crossing_offers_lower_fingerprint_wins() {
1388 let mut rng = TestRng::seeded(8);
1389 let a = Identity::generate(&mut rng).fingerprint();
1390 let b = Identity::generate(&mut rng).fingerprint();
1391 let (lower, higher) = if a < b { (a, b) } else { (b, a) };
1392
1393 assert!(keeps_own_offer(lower, Some(higher)));
1394 assert!(!keeps_own_offer(higher, Some(lower)));
1395 assert!(!keeps_own_offer(lower, None));
1396 }
1397
1398 #[test]
1399 fn changed_fingerprint_is_refused() {
1400 let mut rng = TestRng::seeded(9);
1401 let alice_identity = Identity::generate(&mut rng);
1402 let bob_identity = Identity::generate(&mut rng);
1403 let mallory_identity = Identity::generate(&mut rng);
1404
1405 let mut alice = Session::new();
1406 alice.trust.observe(mallory_identity.fingerprint());
1407
1408 let mut bob = Session::new();
1409 let bundle = alice.start(&alice_identity, &mut rng).unwrap();
1410 bob.receive_offer(bundle);
1411 let response = bob.accept(&bob_identity, &mut rng).unwrap();
1412
1413 let result = alice.receive_accept(&alice_identity, &response, &mut rng);
1414 assert_eq!(
1415 result,
1416 Err(Error::FingerprintChanged {
1417 previous: mallory_identity.fingerprint(),
1418 current: bob_identity.fingerprint(),
1419 })
1420 );
1421 assert!(!alice.is_established());
1422
1423 alice.confirm_fingerprint_change(bob_identity.fingerprint());
1424 alice
1425 .receive_accept(&alice_identity, &response, &mut rng)
1426 .unwrap();
1427 assert!(alice.is_established());
1428 }
1429
1430 #[test]
1431 fn safety_number_is_eight_groups_of_four() {
1432 let mut rng = TestRng::seeded(10);
1433 let fingerprint = Identity::generate(&mut rng).fingerprint();
1434 let rendered = fingerprint.safety_number();
1435 let groups: Vec<&str> = rendered.split(' ').collect();
1436 assert_eq!(groups.len(), 8);
1437 for group in groups {
1438 assert_eq!(group.len(), 4);
1439 assert!(
1440 group
1441 .chars()
1442 .all(|c| c.is_ascii_hexdigit() && !c.is_ascii_lowercase())
1443 );
1444 }
1445 }
1446
1447 #[test]
1448 fn frag_reassembles_in_order_and_rejects_a_gap() {
1449 let whole = b"hello obby world".to_vec();
1450 let (first_half, second_half) = whole.split_at(8);
1451 let fragments = alloc::vec![
1452 Frag {
1453 id: "abc".into(),
1454 i: 0,
1455 n: 2,
1456 ct: first_half.to_vec(),
1457 },
1458 Frag {
1459 id: "abc".into(),
1460 i: 1,
1461 n: 2,
1462 ct: second_half.to_vec(),
1463 },
1464 ];
1465 assert_eq!(reassemble(&fragments).unwrap(), whole);
1466
1467 let missing_second = &fragments[..1];
1468 assert_eq!(reassemble(missing_second), Err(Error::Fragmentation));
1469 }
1470
1471 #[test]
1472 fn reject_and_close_transition_state_and_frame() {
1473 let mut rng = TestRng::seeded(11);
1474 let alice_identity = Identity::generate(&mut rng);
1475 let mut alice = Session::new();
1476 let mut bob = Session::new();
1477
1478 let bundle = alice.start(&alice_identity, &mut rng).unwrap();
1479 let offered = bob.receive_offer(bundle.clone());
1480 assert_eq!(offered, alice_identity.fingerprint());
1481
1482 let frame = bob.reject(Some("busy".into()));
1483 assert_eq!(
1484 frame,
1485 Frame::Reject {
1486 reason: Some("busy".into())
1487 }
1488 );
1489 assert!(!bob.is_established());
1490 alice.receive_reject();
1491 assert_eq!(alice.send(b"too late"), Err(Error::WrongState));
1492
1493 let mut pair = establish();
1494 let frame = pair.alice.close();
1495 assert_eq!(frame, Frame::Close);
1496 pair.bob.receive_close();
1497 assert_eq!(pair.alice.send(b"too late"), Err(Error::WrongState));
1498 assert_eq!(pair.bob.send(b"too late"), Err(Error::WrongState));
1499
1500 let init = Frame::Init {
1502 bundle,
1503 account: Some("alice".into()),
1504 };
1505 assert!(matches!(init, Frame::Init { .. }));
1506 assert_eq!(PROTOCOL_VERSION, 1);
1507 }
1508
1509 #[test]
1510 fn frame_wraps_every_content_and_handshake_variant() {
1511 let mut pair = establish();
1512 let ct = pair.alice.send(b"hi").unwrap();
1513 let msg = Frame::Msg { ct: ct.clone() };
1514 let media = Frame::Media { ct: ct.clone() };
1515 let ack = Frame::Ack { ct };
1516 assert!(matches!(msg, Frame::Msg { .. }));
1517 assert!(matches!(media, Frame::Media { .. }));
1518 assert!(matches!(ack, Frame::Ack { .. }));
1519
1520 let mut rng = TestRng::seeded(12);
1521 let alice_identity = Identity::generate(&mut rng);
1522 let bob_identity = Identity::generate(&mut rng);
1523 let mut alice = Session::new();
1524 let mut bob = Session::new();
1525 let bundle = alice.start(&alice_identity, &mut rng).unwrap();
1526 bob.receive_offer(bundle);
1527 let response = bob.accept(&bob_identity, &mut rng).unwrap();
1528 let accept = Frame::Accept {
1529 response,
1530 account: None,
1531 };
1532 assert!(matches!(accept, Frame::Accept { .. }));
1533 }
1534
1535 #[test]
1536 fn peer_verification_tracks_out_of_band_confirmation() {
1537 let mut pair = establish();
1538 assert!(!pair.alice.is_peer_verified());
1539 pair.alice.mark_peer_verified();
1540 assert!(pair.alice.is_peer_verified());
1541 assert!(!pair.bob.is_peer_verified());
1542 }
1543}