Skip to main content

nostr_double_ratchet/
session.rs

1use crate::{
2    device_pubkey_from_secret_bytes, kdf, random_secret_key_bytes, secret_key_from_bytes,
3    DevicePubkey, DomainError, ProtocolContext, Result, UnixSeconds, MAX_SKIP,
4};
5use base64::Engine;
6use nostr::nips::nip44::{self, Version};
7use nostr::PublicKey;
8use rand::rngs::OsRng;
9use rand::{CryptoRng, RngCore};
10use serde::{Deserialize, Serialize};
11use std::collections::BTreeMap;
12
13#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
14#[serde(rename_all = "camelCase")]
15pub struct Header {
16    pub number: u32,
17    pub previous_chain_length: u32,
18    pub next_public_key: DevicePubkey,
19}
20
21#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
22pub struct SerializableKeyPair {
23    pub public_key: DevicePubkey,
24    #[serde(with = "serde_bytes_array")]
25    pub private_key: [u8; 32],
26}
27
28#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
29pub struct SkippedKeysEntry {
30    #[serde(with = "serde_btreemap_u32_bytes")]
31    pub message_keys: BTreeMap<u32, [u8; 32]>,
32}
33
34#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
35pub struct SessionState {
36    #[serde(with = "serde_bytes_array")]
37    pub root_key: [u8; 32],
38    pub their_current_nostr_public_key: Option<DevicePubkey>,
39    pub their_next_nostr_public_key: Option<DevicePubkey>,
40    #[serde(default, skip_serializing_if = "Option::is_none")]
41    pub our_previous_nostr_key: Option<SerializableKeyPair>,
42    pub our_current_nostr_key: Option<SerializableKeyPair>,
43    pub our_next_nostr_key: SerializableKeyPair,
44    #[serde(default, with = "serde_option_bytes_array")]
45    pub receiving_chain_key: Option<[u8; 32]>,
46    #[serde(default, with = "serde_option_bytes_array")]
47    pub sending_chain_key: Option<[u8; 32]>,
48    pub sending_chain_message_number: u32,
49    pub receiving_chain_message_number: u32,
50    pub previous_sending_chain_message_count: u32,
51    pub skipped_keys: BTreeMap<DevicePubkey, SkippedKeysEntry>,
52}
53
54#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
55pub struct MessageEnvelope {
56    pub sender: DevicePubkey,
57    #[serde(default, skip_serializing_if = "Option::is_none")]
58    pub recipient: Option<DevicePubkey>,
59    pub signer_secret_key: [u8; 32],
60    pub created_at: UnixSeconds,
61    pub encrypted_header: String,
62    pub ciphertext: String,
63}
64
65#[derive(Debug, Clone)]
66pub struct SendPlan {
67    pub next_state: SessionState,
68    pub envelope: MessageEnvelope,
69    pub payload: Vec<u8>,
70}
71
72#[derive(Debug, Clone)]
73pub struct SendOutcome {
74    pub envelope: MessageEnvelope,
75    pub payload: Vec<u8>,
76}
77
78#[derive(Debug, Clone)]
79pub struct ReceivePlan {
80    pub next_state: SessionState,
81    pub payload: Vec<u8>,
82    pub sender: DevicePubkey,
83}
84
85#[derive(Debug, Clone)]
86pub struct ReceiveOutcome {
87    pub payload: Vec<u8>,
88    pub sender: DevicePubkey,
89}
90
91#[derive(Debug, Clone)]
92pub struct Session {
93    pub state: SessionState,
94    pub name: String,
95}
96
97#[derive(Debug, Clone, Copy, PartialEq, Eq)]
98enum HeaderDecryptionTarget {
99    Current,
100    Next,
101    Previous,
102}
103
104impl Session {
105    pub fn from_state(state: SessionState) -> Self {
106        Self {
107            state,
108            name: String::new(),
109        }
110    }
111
112    pub fn new(state: SessionState, name: String) -> Self {
113        Self { state, name }
114    }
115
116    pub fn init(
117        their_ephemeral_nostr_public_key: PublicKey,
118        our_ephemeral_nostr_private_key: [u8; 32],
119        is_initiator: bool,
120        shared_secret: [u8; 32],
121        _name: Option<String>,
122    ) -> Result<Self> {
123        let mut rng = OsRng;
124        let mut ctx = ProtocolContext::new(
125            UnixSeconds(
126                std::time::SystemTime::now()
127                    .duration_since(std::time::UNIX_EPOCH)
128                    .unwrap()
129                    .as_secs(),
130            ),
131            &mut rng,
132        );
133        let peer = DevicePubkey::from_bytes(their_ephemeral_nostr_public_key.to_bytes());
134        let mut session = if is_initiator {
135            Self::new_initiator(
136                &mut ctx,
137                peer,
138                our_ephemeral_nostr_private_key,
139                shared_secret,
140            )
141        } else {
142            Self::new_responder(
143                &mut ctx,
144                peer,
145                our_ephemeral_nostr_private_key,
146                shared_secret,
147            )
148        }?;
149        session.name = _name.unwrap_or_default();
150        Ok(session)
151    }
152
153    pub fn new_initiator<R>(
154        ctx: &mut ProtocolContext<'_, R>,
155        their_ephemeral_public_key: DevicePubkey,
156        our_ephemeral_private_key: [u8; 32],
157        shared_secret: [u8; 32],
158    ) -> Result<Self>
159    where
160        R: RngCore + CryptoRng,
161    {
162        Self::init_with_context(
163            ctx,
164            their_ephemeral_public_key,
165            our_ephemeral_private_key,
166            true,
167            shared_secret,
168        )
169    }
170
171    pub fn new_responder<R>(
172        ctx: &mut ProtocolContext<'_, R>,
173        their_ephemeral_public_key: DevicePubkey,
174        our_ephemeral_private_key: [u8; 32],
175        shared_secret: [u8; 32],
176    ) -> Result<Self>
177    where
178        R: RngCore + CryptoRng,
179    {
180        Self::init_with_context(
181            ctx,
182            their_ephemeral_public_key,
183            our_ephemeral_private_key,
184            false,
185            shared_secret,
186        )
187    }
188
189    fn init_with_context<R>(
190        ctx: &mut ProtocolContext<'_, R>,
191        their_ephemeral_public_key: DevicePubkey,
192        our_ephemeral_private_key: [u8; 32],
193        is_initiator: bool,
194        shared_secret: [u8; 32],
195    ) -> Result<Self>
196    where
197        R: RngCore + CryptoRng,
198    {
199        let our_keys = nostr::Keys::new(secret_key_from_bytes(&our_ephemeral_private_key)?);
200        let our_next_private_key = random_secret_key_bytes(ctx.rng)?;
201        let our_next_keys = nostr::Keys::new(secret_key_from_bytes(&our_next_private_key)?);
202
203        let (root_key, sending_chain_key, our_current_nostr_key, our_next_nostr_key) =
204            if is_initiator {
205                let our_current_pubkey = DevicePubkey::from_nostr(our_keys.public_key());
206                let conversation_key = nip44::v2::ConversationKey::derive(
207                    our_next_keys.secret_key(),
208                    &their_ephemeral_public_key.to_nostr()?,
209                )?;
210                let kdf_outputs = kdf(&shared_secret, conversation_key.as_bytes(), 2);
211                (
212                    kdf_outputs[0],
213                    Some(kdf_outputs[1]),
214                    Some(SerializableKeyPair {
215                        public_key: our_current_pubkey,
216                        private_key: our_ephemeral_private_key,
217                    }),
218                    SerializableKeyPair {
219                        public_key: DevicePubkey::from_nostr(our_next_keys.public_key()),
220                        private_key: our_next_private_key,
221                    },
222                )
223            } else {
224                (
225                    shared_secret,
226                    None,
227                    None,
228                    SerializableKeyPair {
229                        public_key: DevicePubkey::from_nostr(our_keys.public_key()),
230                        private_key: our_ephemeral_private_key,
231                    },
232                )
233            };
234
235        Ok(Self {
236            state: SessionState {
237                root_key,
238                their_current_nostr_public_key: None,
239                their_next_nostr_public_key: Some(their_ephemeral_public_key),
240                our_previous_nostr_key: None,
241                our_current_nostr_key,
242                our_next_nostr_key,
243                receiving_chain_key: None,
244                sending_chain_key,
245                sending_chain_message_number: 0,
246                receiving_chain_message_number: 0,
247                previous_sending_chain_message_count: 0,
248                skipped_keys: BTreeMap::new(),
249            },
250            name: String::new(),
251        })
252    }
253
254    pub fn can_send(&self) -> bool {
255        self.state.their_next_nostr_public_key.is_some()
256            && self.state.our_current_nostr_key.is_some()
257    }
258
259    pub fn matches_sender(&self, sender: DevicePubkey) -> bool {
260        self.state.their_current_nostr_public_key == Some(sender)
261            || self.state.their_next_nostr_public_key == Some(sender)
262            || self.state.skipped_keys.contains_key(&sender)
263    }
264
265    pub fn plan_send(&self, payload: &[u8], now: UnixSeconds) -> Result<SendPlan> {
266        if !self.can_send() {
267            return Err(DomainError::CannotSendYet.into());
268        }
269
270        let mut next_state = self.state.clone();
271        let (header, ciphertext) = ratchet_encrypt(&mut next_state, payload)?;
272        let our_current = self
273            .state
274            .our_current_nostr_key
275            .as_ref()
276            .ok_or(DomainError::SessionNotReady)?;
277        let our_secret = secret_key_from_bytes(&our_current.private_key)?;
278        let their_next = self
279            .state
280            .their_next_nostr_public_key
281            .ok_or(DomainError::SessionNotReady)?;
282        let encrypted_header = nip44::encrypt(
283            &our_secret,
284            &their_next.to_nostr()?,
285            &serde_json::to_string(&header)?,
286            Version::V2,
287        )?;
288
289        Ok(SendPlan {
290            next_state,
291            envelope: MessageEnvelope {
292                sender: our_current.public_key,
293                recipient: None,
294                signer_secret_key: our_current.private_key,
295                created_at: now,
296                encrypted_header,
297                ciphertext,
298            },
299            payload: payload.to_vec(),
300        })
301    }
302
303    pub fn apply_send(&mut self, plan: SendPlan) -> SendOutcome {
304        self.state = plan.next_state;
305        SendOutcome {
306            envelope: plan.envelope,
307            payload: plan.payload,
308        }
309    }
310
311    pub fn plan_receive<R>(
312        &self,
313        ctx: &mut ProtocolContext<'_, R>,
314        envelope: &MessageEnvelope,
315    ) -> Result<ReceivePlan>
316    where
317        R: RngCore + CryptoRng,
318    {
319        if !self.matches_sender(envelope.sender) {
320            return Err(DomainError::UnexpectedSender.into());
321        }
322
323        let mut next_state = self.state.clone();
324        let previous_chain_sender = next_state
325            .their_current_nostr_public_key
326            .or(next_state.their_next_nostr_public_key);
327        let (header, decryption_target) =
328            decrypt_header(&next_state, &envelope.encrypted_header, envelope.sender)?;
329        let should_ratchet = decryption_target == HeaderDecryptionTarget::Next;
330
331        let expected_next = next_state.their_next_nostr_public_key;
332        if should_ratchet && expected_next != Some(header.next_public_key) {
333            next_state.their_current_nostr_public_key = next_state.their_next_nostr_public_key;
334            next_state.their_next_nostr_public_key = Some(header.next_public_key);
335        }
336
337        if should_ratchet {
338            if next_state.receiving_chain_key.is_some() {
339                let skipped_sender = previous_chain_sender.ok_or(DomainError::SessionNotReady)?;
340                skip_message_keys(
341                    &mut next_state,
342                    header.previous_chain_length,
343                    skipped_sender,
344                )?;
345            }
346            ratchet_step(&mut next_state, ctx.rng)?;
347        }
348
349        let payload = ratchet_decrypt(
350            &mut next_state,
351            &header,
352            &envelope.ciphertext,
353            envelope.sender,
354        )?;
355
356        Ok(ReceivePlan {
357            next_state,
358            payload,
359            sender: envelope.sender,
360        })
361    }
362
363    pub fn apply_receive(&mut self, plan: ReceivePlan) -> ReceiveOutcome {
364        self.state = plan.next_state;
365        ReceiveOutcome {
366            payload: plan.payload,
367            sender: plan.sender,
368        }
369    }
370
371    pub fn close(&self) {}
372}
373
374fn ratchet_encrypt(state: &mut SessionState, plaintext: &[u8]) -> Result<(Header, String)> {
375    let sending_chain_key = state
376        .sending_chain_key
377        .ok_or(DomainError::SessionNotReady)?;
378
379    let kdf_outputs = kdf(&sending_chain_key, &[1u8], 2);
380    state.sending_chain_key = Some(kdf_outputs[0]);
381    let message_key = kdf_outputs[1];
382
383    let header = Header {
384        number: state.sending_chain_message_number,
385        next_public_key: state.our_next_nostr_key.public_key,
386        previous_chain_length: state.previous_sending_chain_message_count,
387    };
388
389    state.sending_chain_message_number += 1;
390
391    let conversation_key = nip44::v2::ConversationKey::new(message_key);
392    let encrypted_bytes = nip44::v2::encrypt_to_bytes(&conversation_key, plaintext)?;
393    let ciphertext = base64::engine::general_purpose::STANDARD.encode(encrypted_bytes);
394    Ok((header, ciphertext))
395}
396
397fn ratchet_decrypt(
398    state: &mut SessionState,
399    header: &Header,
400    ciphertext: &str,
401    sender: DevicePubkey,
402) -> Result<Vec<u8>> {
403    if let Some(plaintext) = try_skipped_message_keys(state, header, ciphertext, sender)? {
404        return Ok(plaintext);
405    }
406
407    if state.receiving_chain_key.is_none() {
408        return Err(DomainError::SessionNotReady.into());
409    }
410
411    skip_message_keys(state, header.number, sender)?;
412
413    let receiving_chain_key = state
414        .receiving_chain_key
415        .ok_or(DomainError::SessionNotReady)?;
416
417    let kdf_outputs = kdf(&receiving_chain_key, &[1u8], 2);
418    state.receiving_chain_key = Some(kdf_outputs[0]);
419    let message_key = kdf_outputs[1];
420    state.receiving_chain_message_number += 1;
421
422    let conversation_key = nip44::v2::ConversationKey::new(message_key);
423    let ciphertext_bytes = base64::engine::general_purpose::STANDARD
424        .decode(ciphertext)
425        .map_err(|e| crate::Error::Decryption(e.to_string()))?;
426
427    nip44::v2::decrypt_to_bytes(&conversation_key, &ciphertext_bytes).map_err(Into::into)
428}
429
430fn ratchet_step<R>(state: &mut SessionState, rng: &mut R) -> Result<()>
431where
432    R: RngCore + CryptoRng,
433{
434    state.previous_sending_chain_message_count = state.sending_chain_message_number;
435    state.sending_chain_message_number = 0;
436    state.receiving_chain_message_number = 0;
437
438    let our_next_sk = secret_key_from_bytes(&state.our_next_nostr_key.private_key)?;
439    let their_next_pk = state
440        .their_next_nostr_public_key
441        .ok_or(DomainError::SessionNotReady)?;
442
443    let conversation_key1 =
444        nip44::v2::ConversationKey::derive(&our_next_sk, &their_next_pk.to_nostr()?)?;
445    let kdf_outputs = kdf(&state.root_key, conversation_key1.as_bytes(), 2);
446    state.receiving_chain_key = Some(kdf_outputs[1]);
447    state.our_previous_nostr_key = state.our_current_nostr_key.clone();
448    state.our_current_nostr_key = Some(state.our_next_nostr_key.clone());
449
450    let our_next_private_key = random_secret_key_bytes(rng)?;
451    state.our_next_nostr_key = SerializableKeyPair {
452        public_key: device_pubkey_from_secret_bytes(&our_next_private_key)?,
453        private_key: our_next_private_key,
454    };
455
456    let our_next_sk2 = secret_key_from_bytes(&our_next_private_key)?;
457    let conversation_key2 =
458        nip44::v2::ConversationKey::derive(&our_next_sk2, &their_next_pk.to_nostr()?)?;
459    let kdf_outputs2 = kdf(&kdf_outputs[0], conversation_key2.as_bytes(), 2);
460    state.root_key = kdf_outputs2[0];
461    state.sending_chain_key = Some(kdf_outputs2[1]);
462    Ok(())
463}
464
465fn skip_message_keys(state: &mut SessionState, until: u32, sender: DevicePubkey) -> Result<()> {
466    if until <= state.receiving_chain_message_number {
467        return Ok(());
468    }
469
470    if (until - state.receiving_chain_message_number) as usize > MAX_SKIP {
471        return Err(DomainError::TooManySkippedMessages.into());
472    }
473
474    let entry = state.skipped_keys.entry(sender).or_default();
475
476    while state.receiving_chain_message_number < until {
477        let receiving_chain_key = state
478            .receiving_chain_key
479            .ok_or(DomainError::SessionNotReady)?;
480        let kdf_outputs = kdf(&receiving_chain_key, &[1u8], 2);
481        state.receiving_chain_key = Some(kdf_outputs[0]);
482        entry
483            .message_keys
484            .insert(state.receiving_chain_message_number, kdf_outputs[1]);
485        state.receiving_chain_message_number += 1;
486    }
487
488    prune_skipped_message_keys(&mut entry.message_keys);
489    Ok(())
490}
491
492fn try_skipped_message_keys(
493    state: &mut SessionState,
494    header: &Header,
495    ciphertext: &str,
496    sender: DevicePubkey,
497) -> Result<Option<Vec<u8>>> {
498    if let Some(entry) = state.skipped_keys.get_mut(&sender) {
499        if let Some(message_key) = entry.message_keys.remove(&header.number) {
500            let conversation_key = nip44::v2::ConversationKey::new(message_key);
501            let ciphertext_bytes = base64::engine::general_purpose::STANDARD
502                .decode(ciphertext)
503                .map_err(|e| crate::Error::Decryption(e.to_string()))?;
504            let plaintext = nip44::v2::decrypt_to_bytes(&conversation_key, &ciphertext_bytes)?;
505            if entry.message_keys.is_empty() {
506                state.skipped_keys.remove(&sender);
507            }
508            return Ok(Some(plaintext));
509        }
510    }
511
512    Ok(None)
513}
514
515fn decrypt_header(
516    state: &SessionState,
517    encrypted_header: &str,
518    sender: DevicePubkey,
519) -> Result<(Header, HeaderDecryptionTarget)> {
520    if let Some(current) = &state.our_current_nostr_key {
521        let current_sk = secret_key_from_bytes(&current.private_key)?;
522        if let Ok(decrypted) = nip44::decrypt(&current_sk, &sender.to_nostr()?, encrypted_header) {
523            let header: Header = serde_json::from_str(&decrypted)?;
524            return Ok((header, HeaderDecryptionTarget::Current));
525        }
526    }
527
528    let next_sk = secret_key_from_bytes(&state.our_next_nostr_key.private_key)?;
529    if let Ok(decrypted) = nip44::decrypt(&next_sk, &sender.to_nostr()?, encrypted_header) {
530        let header: Header = serde_json::from_str(&decrypted)?;
531        return Ok((header, HeaderDecryptionTarget::Next));
532    }
533
534    if let Some(previous) = &state.our_previous_nostr_key {
535        let previous_sk = secret_key_from_bytes(&previous.private_key)?;
536        if let Ok(decrypted) = nip44::decrypt(&previous_sk, &sender.to_nostr()?, encrypted_header) {
537            let header: Header = serde_json::from_str(&decrypted)?;
538            return Ok((header, HeaderDecryptionTarget::Previous));
539        }
540    }
541
542    Err(crate::Error::Parse("invalid header".to_string()))
543}
544
545fn prune_skipped_message_keys(map: &mut BTreeMap<u32, [u8; 32]>) {
546    while map.len() > MAX_SKIP {
547        let Some(first) = map.keys().next().copied() else {
548            break;
549        };
550        map.remove(&first);
551    }
552}
553
554mod serde_bytes_array {
555    use serde::{Deserialize, Deserializer, Serializer};
556
557    pub fn serialize<S>(bytes: &[u8; 32], serializer: S) -> Result<S::Ok, S::Error>
558    where
559        S: Serializer,
560    {
561        serializer.serialize_str(&hex::encode(bytes))
562    }
563
564    pub fn deserialize<'de, D>(deserializer: D) -> Result<[u8; 32], D::Error>
565    where
566        D: Deserializer<'de>,
567    {
568        let s = String::deserialize(deserializer)?;
569        super::decode_hex_32(&s).map_err(serde::de::Error::custom)
570    }
571}
572
573mod serde_option_bytes_array {
574    use serde::{Deserialize, Deserializer, Serializer};
575
576    pub fn serialize<S>(bytes: &Option<[u8; 32]>, serializer: S) -> Result<S::Ok, S::Error>
577    where
578        S: Serializer,
579    {
580        match bytes {
581            Some(b) => serializer.serialize_str(&hex::encode(b)),
582            None => serializer.serialize_none(),
583        }
584    }
585
586    pub fn deserialize<'de, D>(deserializer: D) -> Result<Option<[u8; 32]>, D::Error>
587    where
588        D: Deserializer<'de>,
589    {
590        let opt: Option<String> = Option::deserialize(deserializer)?;
591        match opt {
592            Some(s) => super::decode_hex_32(&s)
593                .map(Some)
594                .map_err(serde::de::Error::custom),
595            None => Ok(None),
596        }
597    }
598}
599
600mod serde_btreemap_u32_bytes {
601    use serde::{Deserialize, Deserializer, Serialize, Serializer};
602    use std::collections::BTreeMap;
603
604    pub fn serialize<S>(map: &BTreeMap<u32, [u8; 32]>, serializer: S) -> Result<S::Ok, S::Error>
605    where
606        S: Serializer,
607    {
608        let string_map: BTreeMap<String, String> = map
609            .iter()
610            .map(|(k, v)| (k.to_string(), hex::encode(v)))
611            .collect();
612        string_map.serialize(serializer)
613    }
614
615    pub fn deserialize<'de, D>(deserializer: D) -> Result<BTreeMap<u32, [u8; 32]>, D::Error>
616    where
617        D: Deserializer<'de>,
618    {
619        let string_map: BTreeMap<String, String> = BTreeMap::deserialize(deserializer)?;
620        let mut out = BTreeMap::new();
621        for (k, v) in string_map {
622            let idx: u32 = k.parse().map_err(serde::de::Error::custom)?;
623            out.insert(
624                idx,
625                super::decode_hex_32(&v).map_err(serde::de::Error::custom)?,
626            );
627        }
628        Ok(out)
629    }
630}
631
632fn decode_hex_32(value: &str) -> std::result::Result<[u8; 32], String> {
633    let bytes = hex::decode(value).map_err(|e| e.to_string())?;
634    <[u8; 32]>::try_from(bytes.as_slice()).map_err(|_| "invalid 32-byte hex".to_string())
635}
636
637#[cfg(test)]
638mod tests {
639    use super::*;
640    use rand::{rngs::StdRng, SeedableRng};
641
642    fn context(seed: u64) -> ProtocolContext<'static, StdRng> {
643        let rng = Box::new(StdRng::seed_from_u64(seed));
644        let rng = Box::leak(rng);
645        ProtocolContext::new(UnixSeconds(1_700_000_000), rng)
646    }
647
648    #[test]
649    fn header_json_uses_camel_case_wire_fields() {
650        let header = Header {
651            number: 3,
652            previous_chain_length: 2,
653            next_public_key: DevicePubkey::from_bytes([9u8; 32]),
654        };
655
656        let json = serde_json::to_value(&header).unwrap();
657        assert_eq!(json["number"], serde_json::json!(3));
658        assert_eq!(json["previousChainLength"], serde_json::json!(2));
659        assert_eq!(
660            json["nextPublicKey"],
661            serde_json::json!(header.next_public_key.to_string())
662        );
663        assert!(json.get("previous_chain_length").is_none());
664        assert!(json.get("next_public_key").is_none());
665
666        let decoded: Header = serde_json::from_value(json).unwrap();
667        assert_eq!(decoded, header);
668    }
669
670    #[test]
671    fn header_json_rejects_snake_case_wire_fields() {
672        let old_header = serde_json::json!({
673            "number": 3,
674            "previous_chain_length": 2,
675            "next_public_key": DevicePubkey::from_bytes([9u8; 32]).to_string(),
676        });
677
678        assert!(serde_json::from_value::<Header>(old_header).is_err());
679    }
680
681    #[test]
682    fn plan_send_and_apply_receive_roundtrip() {
683        let alice_secret = [1u8; 32];
684        let bob_secret = [2u8; 32];
685        let alice_pub = device_pubkey_from_secret_bytes(&alice_secret).unwrap();
686        let bob_pub = device_pubkey_from_secret_bytes(&bob_secret).unwrap();
687        let shared_secret = [7u8; 32];
688
689        let mut init_ctx_alice = context(1);
690        let alice =
691            Session::new_initiator(&mut init_ctx_alice, bob_pub, alice_secret, shared_secret)
692                .unwrap();
693        let mut init_ctx_bob = context(2);
694        let mut bob =
695            Session::new_responder(&mut init_ctx_bob, alice_pub, bob_secret, shared_secret)
696                .unwrap();
697
698        let payload = b"hello".to_vec();
699        let send_plan = alice
700            .plan_send(&payload, UnixSeconds(1_700_000_010))
701            .unwrap();
702        let send_outcome = alice.clone().apply_send(send_plan.clone());
703
704        let mut recv_ctx = context(10);
705        let receive_plan = bob
706            .plan_receive(&mut recv_ctx, &send_outcome.envelope)
707            .unwrap();
708        let outcome = bob.apply_receive(receive_plan);
709        assert_eq!(outcome.payload, payload);
710    }
711
712    #[test]
713    fn plan_receive_does_not_mutate_original_session() {
714        let alice_secret = [3u8; 32];
715        let bob_secret = [4u8; 32];
716        let alice_pub = device_pubkey_from_secret_bytes(&alice_secret).unwrap();
717        let bob_pub = device_pubkey_from_secret_bytes(&bob_secret).unwrap();
718        let shared_secret = [8u8; 32];
719
720        let mut init_ctx_alice = context(3);
721        let alice =
722            Session::new_initiator(&mut init_ctx_alice, bob_pub, alice_secret, shared_secret)
723                .unwrap();
724        let mut init_ctx_bob = context(4);
725        let bob = Session::new_responder(&mut init_ctx_bob, alice_pub, bob_secret, shared_secret)
726            .unwrap();
727        let bob_before = bob.state.clone();
728
729        let payload = b"typing".to_vec();
730        let send_plan = alice
731            .plan_send(&payload, UnixSeconds(1_700_000_011))
732            .unwrap();
733
734        let mut recv_ctx = context(13);
735        let _ = bob
736            .plan_receive(&mut recv_ctx, &send_plan.envelope)
737            .unwrap();
738
739        assert_eq!(bob.state, bob_before);
740    }
741
742    #[test]
743    fn duplicate_receive_fails_without_corrupting_state() {
744        let alice_secret = [5u8; 32];
745        let bob_secret = [6u8; 32];
746        let alice_pub = device_pubkey_from_secret_bytes(&alice_secret).unwrap();
747        let bob_pub = device_pubkey_from_secret_bytes(&bob_secret).unwrap();
748        let shared_secret = [9u8; 32];
749
750        let mut init_ctx_alice = context(5);
751        let alice =
752            Session::new_initiator(&mut init_ctx_alice, bob_pub, alice_secret, shared_secret)
753                .unwrap();
754        let mut init_ctx_bob = context(6);
755        let mut bob =
756            Session::new_responder(&mut init_ctx_bob, alice_pub, bob_secret, shared_secret)
757                .unwrap();
758
759        let payload = b"hello".to_vec();
760        let send_plan = alice
761            .plan_send(&payload, UnixSeconds(1_700_000_012))
762            .unwrap();
763        let envelope = alice.clone().apply_send(send_plan).envelope;
764
765        let mut recv_ctx = context(15);
766        let first_plan = bob.plan_receive(&mut recv_ctx, &envelope).unwrap();
767        let _ = bob.apply_receive(first_plan);
768        let after_first = bob.state.clone();
769
770        let mut replay_ctx = context(16);
771        let replay = bob.plan_receive(&mut replay_ctx, &envelope);
772        assert!(replay.is_err());
773        assert_eq!(bob.state, after_first);
774    }
775
776    #[test]
777    fn invalid_sender_is_rejected() {
778        let alice_secret = [7u8; 32];
779        let bob_secret = [8u8; 32];
780        let alice_pub = device_pubkey_from_secret_bytes(&alice_secret).unwrap();
781        let shared_secret = [10u8; 32];
782
783        let mut init_ctx_bob = context(7);
784        let bob = Session::new_responder(&mut init_ctx_bob, alice_pub, bob_secret, shared_secret)
785            .unwrap();
786
787        let mut recv_ctx = context(17);
788        let err = bob
789            .plan_receive(
790                &mut recv_ctx,
791                &MessageEnvelope {
792                    sender: device_pubkey_from_secret_bytes(&bob_secret).unwrap(),
793                    recipient: None,
794                    signer_secret_key: bob_secret,
795                    created_at: UnixSeconds(1),
796                    encrypted_header: "bad".to_string(),
797                    ciphertext: "bad".to_string(),
798                },
799            )
800            .unwrap_err();
801        assert!(matches!(
802            err,
803            crate::Error::Domain(DomainError::UnexpectedSender)
804        ));
805    }
806}