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(¤t.private_key)?;
522 if let Ok(decrypted) = nip44::decrypt(¤t_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}