1use crate::store::Device;
2use async_lock::Mutex;
3use async_trait::async_trait;
4use rand::RngExt;
5use std::sync::{Arc, OnceLock};
6use wacore::libsignal::protocol::error::Result as SignalResult;
7use wacore::libsignal::protocol::{
8 Direction, IdentityChange, IdentityKey, IdentityKeyPair, IdentityKeyStore, PrivateKey,
9 ProtocolAddress, PublicKey, SenderKeyRecord, SenderKeyStore, SessionRecord,
10 SignalProtocolError,
11};
12use wacore::libsignal::store::sender_key_name::SenderKeyName;
13use wacore::libsignal::store::*;
14use waproto::whatsapp::{PreKeyRecordStructure, SignedPreKeyRecordStructure};
15
16type StoreError = Box<dyn std::error::Error + Send + Sync>;
17
18type DirectStoreIncarnation = [u8; 16];
19
20fn direct_store_incarnation() -> &'static DirectStoreIncarnation {
22 static INCARNATION: OnceLock<DirectStoreIncarnation> = OnceLock::new();
23 INCARNATION.get_or_init(|| {
24 let mut incarnation = [0; 16];
25 rand::make_rng::<rand::rngs::StdRng>().fill(&mut incarnation);
26 incarnation
27 })
28}
29
30macro_rules! impl_store_wrapper {
31 ($wrapper_ty:ty, $read_lock:ident, $write_lock:ident) => {
32 #[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
33 #[cfg_attr(not(target_arch = "wasm32"), async_trait)]
34 impl IdentityKeyStore for $wrapper_ty {
35 async fn get_identity_key_pair(&self) -> SignalResult<IdentityKeyPair> {
36 self.0.$read_lock().await.get_identity_key_pair().await
37 }
38
39 async fn get_local_registration_id(&self) -> SignalResult<u32> {
40 self.0.$read_lock().await.get_local_registration_id().await
41 }
42
43 async fn save_identity(
44 &mut self,
45 address: &ProtocolAddress,
46 identity_key: &IdentityKey,
47 ) -> SignalResult<IdentityChange> {
48 self.0
49 .$write_lock()
50 .await
51 .save_identity(address, identity_key)
52 .await
53 }
54
55 async fn is_trusted_identity(
56 &self,
57 address: &ProtocolAddress,
58 identity_key: &IdentityKey,
59 direction: Direction,
60 ) -> SignalResult<bool> {
61 self.0
62 .$read_lock()
63 .await
64 .is_trusted_identity(address, identity_key, direction)
65 .await
66 }
67
68 async fn get_identity(
69 &self,
70 address: &ProtocolAddress,
71 ) -> SignalResult<Option<IdentityKey>> {
72 self.0.$read_lock().await.get_identity(address).await
73 }
74 }
75
76 #[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
77 #[cfg_attr(not(target_arch = "wasm32"), async_trait)]
78 impl PreKeyStore for $wrapper_ty {
79 async fn load_prekey(
80 &self,
81 prekey_id: u32,
82 ) -> Result<Option<PreKeyRecordStructure>, StoreError> {
83 self.0.$read_lock().await.load_prekey(prekey_id).await
84 }
85
86 async fn store_prekey(
87 &self,
88 prekey_id: u32,
89 record: PreKeyRecordStructure,
90 uploaded: bool,
91 ) -> Result<(), StoreError> {
92 self.0
93 .$write_lock()
94 .await
95 .store_prekey(prekey_id, record, uploaded)
96 .await
97 }
98
99 async fn contains_prekey(&self, prekey_id: u32) -> Result<bool, StoreError> {
100 self.0.$read_lock().await.contains_prekey(prekey_id).await
101 }
102
103 async fn remove_prekey(&self, prekey_id: u32) -> Result<(), StoreError> {
104 self.0.$write_lock().await.remove_prekey(prekey_id).await
105 }
106 }
107
108 #[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
109 #[cfg_attr(not(target_arch = "wasm32"), async_trait)]
110 impl SignedPreKeyStore for $wrapper_ty {
111 async fn load_signed_prekey(
112 &self,
113 signed_prekey_id: u32,
114 ) -> Result<Option<SignedPreKeyRecordStructure>, StoreError> {
115 self.0
116 .$read_lock()
117 .await
118 .load_signed_prekey(signed_prekey_id)
119 .await
120 }
121
122 async fn load_signed_prekeys(
123 &self,
124 ) -> Result<Vec<SignedPreKeyRecordStructure>, StoreError> {
125 self.0.$read_lock().await.load_signed_prekeys().await
126 }
127
128 async fn store_signed_prekey(
129 &self,
130 signed_prekey_id: u32,
131 record: SignedPreKeyRecordStructure,
132 ) -> Result<(), StoreError> {
133 self.0
134 .$write_lock()
135 .await
136 .store_signed_prekey(signed_prekey_id, record)
137 .await
138 }
139
140 async fn contains_signed_prekey(
141 &self,
142 signed_prekey_id: u32,
143 ) -> Result<bool, StoreError> {
144 self.0
145 .$read_lock()
146 .await
147 .contains_signed_prekey(signed_prekey_id)
148 .await
149 }
150
151 async fn remove_signed_prekey(&self, signed_prekey_id: u32) -> Result<(), StoreError> {
152 self.0
153 .$write_lock()
154 .await
155 .remove_signed_prekey(signed_prekey_id)
156 .await
157 }
158 }
159
160 #[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
161 #[cfg_attr(not(target_arch = "wasm32"), async_trait)]
162 impl SessionStore for $wrapper_ty {
163 async fn load_session(
164 &self,
165 address: &ProtocolAddress,
166 ) -> Result<SessionRecord, StoreError> {
167 self.0.$read_lock().await.load_session(address).await
168 }
169
170 async fn get_sub_device_sessions(&self, name: &str) -> Result<Vec<u32>, StoreError> {
171 self.0
172 .$read_lock()
173 .await
174 .get_sub_device_sessions(name)
175 .await
176 }
177
178 async fn store_session(
179 &self,
180 address: &ProtocolAddress,
181 record: &SessionRecord,
182 ) -> Result<(), StoreError> {
183 self.0
184 .$write_lock()
185 .await
186 .store_session(address, record)
187 .await
188 }
189
190 async fn contains_session(
191 &self,
192 address: &ProtocolAddress,
193 ) -> Result<bool, StoreError> {
194 self.0.$read_lock().await.contains_session(address).await
195 }
196
197 async fn delete_session(&self, address: &ProtocolAddress) -> Result<(), StoreError> {
198 self.0.$write_lock().await.delete_session(address).await
199 }
200
201 async fn delete_all_sessions(&self, name: &str) -> Result<(), StoreError> {
202 self.0.$write_lock().await.delete_all_sessions(name).await
203 }
204 }
205 };
206}
207
208#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
209#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
210impl IdentityKeyStore for Device {
211 async fn get_identity_key_pair(&self) -> SignalResult<IdentityKeyPair> {
212 Ok(self.identity_key.clone().into())
213 }
214
215 async fn get_local_registration_id(&self) -> SignalResult<u32> {
216 Ok(self.registration_id)
217 }
218
219 async fn save_identity(
220 &mut self,
221 address: &ProtocolAddress,
222 identity_key: &IdentityKey,
223 ) -> SignalResult<IdentityChange> {
224 let address_str = address.as_str();
225 let key_bytes = identity_key.public_key().public_key_bytes();
226 let existing_identity_opt = self.get_identity(address).await?;
227
228 self.backend
229 .put_identity(
230 address_str,
231 key_bytes.try_into().map_err(|_| {
232 SignalProtocolError::InvalidArgument("Invalid key length".into())
233 })?,
234 )
235 .await
236 .map_err(|e| SignalProtocolError::BackendError("backend put_identity", Box::new(e)))?;
237
238 match existing_identity_opt {
239 None => Ok(IdentityChange::NewOrUnchanged),
240 Some(existing) if &existing == identity_key => Ok(IdentityChange::NewOrUnchanged),
241 Some(_) => Ok(IdentityChange::ReplacedExisting),
242 }
243 }
244
245 async fn is_trusted_identity(
246 &self,
247 _address: &ProtocolAddress,
248 _identity_key: &IdentityKey,
249 _direction: Direction,
250 ) -> SignalResult<bool> {
251 Ok(true)
255 }
256
257 async fn get_identity(&self, address: &ProtocolAddress) -> SignalResult<Option<IdentityKey>> {
258 let identity_bytes = self
259 .backend
260 .load_identity(address.as_str())
261 .await
262 .map_err(|e| SignalProtocolError::BackendError("backend get_identity", Box::new(e)))?;
263
264 match identity_bytes {
265 Some(bytes) if !bytes.is_empty() => {
266 let public_key = PublicKey::from_djb_public_key_bytes(&bytes)?;
267 Ok(Some(IdentityKey::new(public_key)))
268 }
269 _ => Ok(None),
270 }
271 }
272}
273
274#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
275#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
276impl PreKeyStore for Device {
277 async fn load_prekey(
278 &self,
279 prekey_id: u32,
280 ) -> Result<Option<PreKeyRecordStructure>, StoreError> {
281 use wacore::libsignal::protocol::KeyPair;
282 use wacore::libsignal::store::record_helpers::new_pre_key_record;
283
284 match self.backend.load_prekey(prekey_id).await {
285 Ok(Some(bytes)) => {
286 if let Ok(record) = waproto::codec::pre_key_record_decode(&bytes) {
288 return Ok(Some(record));
289 }
290
291 if let Ok(private_key) = PrivateKey::deserialize(&bytes)
294 && let Ok(public_key) = private_key.public_key()
295 {
296 let key_pair = KeyPair::new(public_key, private_key);
297 let record = new_pre_key_record(prekey_id, &key_pair);
298 return Ok(Some(record));
299 }
300
301 Ok(None)
303 }
304 Ok(None) => Ok(None),
305 Err(e) => Err(Box::new(e) as StoreError),
306 }
307 }
308
309 async fn store_prekey(
310 &self,
311 prekey_id: u32,
312 record: PreKeyRecordStructure,
313 uploaded: bool,
314 ) -> Result<(), StoreError> {
315 let bytes = waproto::codec::pre_key_record_to_vec(&record);
316 self.backend
317 .store_prekey(prekey_id, &bytes, uploaded)
318 .await
319 .map_err(|e| Box::new(e) as StoreError)
320 }
321
322 async fn contains_prekey(&self, prekey_id: u32) -> Result<bool, StoreError> {
323 match self.backend.load_prekey(prekey_id).await {
324 Ok(opt) => Ok(opt.is_some()),
325 Err(e) => Err(Box::new(e) as StoreError),
326 }
327 }
328
329 async fn remove_prekey(&self, prekey_id: u32) -> Result<(), StoreError> {
330 self.backend
331 .remove_prekey(prekey_id)
332 .await
333 .map_err(|e| Box::new(e) as StoreError)
334 }
335}
336
337#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
338#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
339impl SignedPreKeyStore for Device {
340 async fn load_signed_prekey(
341 &self,
342 signed_prekey_id: u32,
343 ) -> Result<Option<SignedPreKeyRecordStructure>, StoreError> {
344 if signed_prekey_id == self.signed_pre_key_id {
345 let record = record_helpers::new_signed_pre_key_record(
346 self.signed_pre_key_id,
347 &self.signed_pre_key,
348 self.signed_pre_key_signature,
349 wacore::time::now_utc(),
350 );
351 return Ok(Some(record));
352 }
353 match self
356 .backend
357 .load_signed_prekey(signed_prekey_id)
358 .await
359 .map_err(|e| Box::new(e) as StoreError)?
360 {
361 Some(bytes) => {
362 let record = waproto::codec::signed_pre_key_record_decode(&bytes)
363 .map_err(|e| Box::new(e) as StoreError)?;
364 Ok(Some(record))
365 }
366 None => Ok(None),
367 }
368 }
369
370 async fn load_signed_prekeys(&self) -> Result<Vec<SignedPreKeyRecordStructure>, StoreError> {
371 log::warn!(
372 "Device: load_signed_prekeys() - returning empty list. Only the device's own signed pre-key should be accessed via load_signed_prekey()."
373 );
374 Ok(Vec::new())
375 }
376
377 async fn store_signed_prekey(
378 &self,
379 signed_prekey_id: u32,
380 _record: SignedPreKeyRecordStructure,
381 ) -> Result<(), StoreError> {
382 log::warn!(
383 "Device: store_signed_prekey({}) - no-op. Signed pre-keys should only be set once during device creation/pairing and managed via PersistenceManager.",
384 signed_prekey_id
385 );
386 Ok(())
387 }
388
389 async fn contains_signed_prekey(&self, signed_prekey_id: u32) -> Result<bool, StoreError> {
390 if signed_prekey_id == self.signed_pre_key_id {
391 return Ok(true);
392 }
393 Ok(self
397 .backend
398 .load_signed_prekey(signed_prekey_id)
399 .await
400 .map_err(|e| Box::new(e) as StoreError)?
401 .is_some())
402 }
403
404 async fn remove_signed_prekey(&self, signed_prekey_id: u32) -> Result<(), StoreError> {
405 log::warn!(
406 "Device: remove_signed_prekey({}) - no-op. Signed pre-keys are managed via PersistenceManager and should not be removed individually.",
407 signed_prekey_id
408 );
409 Ok(())
410 }
411}
412
413#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
414#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
415impl SessionStore for Device {
416 async fn load_session(&self, address: &ProtocolAddress) -> Result<SessionRecord, StoreError> {
417 let address_str = address.as_str();
418 match self.backend.get_session(address_str).await {
419 Ok(Some(session_data)) => {
420 SessionRecord::deserialize_for_store(&session_data, direct_store_incarnation())
421 .map_err(|e| Box::new(e) as StoreError)
422 }
423 Ok(None) => Ok(SessionRecord::new_fresh()),
424 Err(e) => Err(Box::new(e) as StoreError),
425 }
426 }
427
428 async fn get_sub_device_sessions(&self, name: &str) -> Result<Vec<u32>, StoreError> {
429 let _ = name;
430 Ok(Vec::new())
431 }
432
433 async fn store_session(
434 &self,
435 address: &ProtocolAddress,
436 record: &SessionRecord,
437 ) -> Result<(), StoreError> {
438 let address_str = address.as_str();
439 let mut session_data = Vec::new();
440 record.serialize_into_for_store(&mut session_data, direct_store_incarnation());
441
442 self.backend
443 .put_session(address_str, &session_data)
444 .await
445 .map_err(|e| Box::new(e) as StoreError)
446 }
447
448 async fn contains_session(&self, address: &ProtocolAddress) -> Result<bool, StoreError> {
449 let address_str = address.as_str();
450 self.backend
451 .has_session(address_str)
452 .await
453 .map_err(|e| Box::new(e) as StoreError)
454 }
455
456 async fn delete_session(&self, address: &ProtocolAddress) -> Result<(), StoreError> {
457 let address_str = address.as_str();
458 self.backend
459 .delete_session(address_str)
460 .await
461 .map_err(|e| Box::new(e) as StoreError)
462 }
463
464 async fn delete_all_sessions(&self, name: &str) -> Result<(), StoreError> {
465 let _ = name;
466 Ok(())
467 }
468}
469
470use async_lock::RwLock;
471
472pub struct DeviceRwLockWrapper(pub Arc<RwLock<Device>>);
473
474impl DeviceRwLockWrapper {
475 pub fn new(device: Arc<RwLock<Device>>) -> Self {
476 Self(device)
477 }
478}
479
480impl Clone for DeviceRwLockWrapper {
481 fn clone(&self) -> Self {
482 Self(self.0.clone())
483 }
484}
485
486impl_store_wrapper!(DeviceRwLockWrapper, read, write);
487
488pub struct DeviceStore(pub Arc<Mutex<Device>>);
489
490impl DeviceStore {
491 pub fn new(device: Arc<Mutex<Device>>) -> Self {
492 Self(device)
493 }
494}
495
496impl Clone for DeviceStore {
497 fn clone(&self) -> Self {
498 Self(self.0.clone())
499 }
500}
501
502impl_store_wrapper!(DeviceStore, lock, lock);
503
504#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
505#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
506impl SenderKeyStore for Device {
507 async fn store_sender_key(
508 &mut self,
509 sender_key_name: &SenderKeyName,
510 record: SenderKeyRecord,
511 ) -> SignalResult<()> {
512 let serialized_record = record.serialize_for_store(direct_store_incarnation())?;
513 self.backend
514 .put_sender_key(sender_key_name.cache_key(), &serialized_record)
515 .await
516 .map_err(|e| SignalProtocolError::BackendError("store_sender_key", Box::new(e)))
517 }
518
519 async fn load_sender_key(
520 &self,
521 sender_key_name: &SenderKeyName,
522 ) -> SignalResult<Option<SenderKeyRecord>> {
523 match self
524 .backend
525 .get_sender_key(sender_key_name.cache_key())
526 .await
527 .map_err(|e| SignalProtocolError::BackendError("load_sender_key", Box::new(e)))?
528 {
529 Some(data) => {
530 let record =
531 SenderKeyRecord::deserialize_for_store(&data, direct_store_incarnation())?;
532 if record.is_empty() {
533 Ok(None)
534 } else {
535 Ok(Some(record))
536 }
537 }
538 None => Ok(None),
539 }
540 }
541}
542
543#[cfg(test)]
544#[allow(clippy::disallowed_methods)]
545mod tests {
546 use super::*;
547
548 fn leased_session() -> SessionRecord {
549 use wacore::libsignal::protocol::{ChainKey, KeyPair, RootKey, SessionState};
550
551 let mut rng = rand::make_rng::<rand::rngs::StdRng>();
552 let local = IdentityKey::new(KeyPair::generate(&mut rng).public_key);
553 let remote = IdentityKey::new(KeyPair::generate(&mut rng).public_key);
554 let base_key = KeyPair::generate(&mut rng).public_key;
555 let mut state = SessionState::new(3, &local, &remote, &RootKey::new([0; 32]), &base_key);
556 state.set_sender_chain(&KeyPair::generate(&mut rng), &ChainKey::new([1; 32], 0));
557 let mut record = SessionRecord::new(state);
558 record.reserve_sender_chain_counters(0);
559 record
560 }
561
562 fn session_chain_index(record: &SessionRecord) -> u32 {
563 record
564 .session_state()
565 .expect("session")
566 .get_sender_chain_key()
567 .expect("sender chain")
568 .index()
569 }
570
571 #[tokio::test]
572 async fn direct_session_store_preserves_clean_reload_and_recovery_ceiling() {
573 let backend = crate::test_utils::create_test_backend().await;
574 let device = Device::new(backend.clone());
575 let address = ProtocolAddress::new("15550001001", 1.into());
576
577 SessionStore::store_session(&device, &address, &leased_session())
578 .await
579 .expect("store session");
580 let clean = SessionStore::load_session(&device, &address)
581 .await
582 .expect("clean reload");
583 assert_eq!(session_chain_index(&clean), 0);
584
585 let replacement = Device::new(backend.clone());
586 let same_process = SessionStore::load_session(&replacement, &address)
587 .await
588 .expect("same-process reload");
589 assert_eq!(session_chain_index(&same_process), 0);
590
591 let durable = backend
592 .get_session(address.as_str())
593 .await
594 .expect("read durable session")
595 .expect("durable session");
596 let recovered = SessionRecord::deserialize(&durable).expect("recovery reload");
597 assert_eq!(
598 session_chain_index(&recovered),
599 wacore::libsignal::protocol::consts::SENDER_CHAIN_RESERVATION_BATCH
600 );
601 }
602
603 #[tokio::test]
604 async fn direct_sender_key_store_preserves_clean_reloads_and_recovery_ceiling() {
605 use wacore::libsignal::protocol::{
606 create_sender_key_distribution_message, group_decrypt, group_encrypt,
607 process_sender_key_distribution_message,
608 };
609
610 let sender_backend = crate::test_utils::create_test_backend().await;
611 let mut sender = Device::new(sender_backend.clone());
612 let mut receiver = Device::new(crate::test_utils::create_test_backend().await);
613 let name = SenderKeyName::from_parts("1234567890@g.us", "15550001000@s.whatsapp.net:0");
614 let mut rng = rand::make_rng::<rand::rngs::StdRng>();
615 let distribution = create_sender_key_distribution_message(&name, &mut sender, &mut rng)
616 .await
617 .expect("sender setup");
618 process_sender_key_distribution_message(&name, &distribution, &mut receiver)
619 .await
620 .expect("receiver setup");
621
622 let mut last = None;
623 for expected_iteration in 0..=32 {
624 let message = group_encrypt(&mut sender, &name, b"payload", &mut rng)
625 .await
626 .expect("group encrypt");
627 assert_eq!(message.iteration(), expected_iteration);
628 last = Some(message);
629 }
630
631 let plaintext = group_decrypt(
632 last.expect("last message").serialized(),
633 &mut receiver,
634 &name,
635 )
636 .await
637 .expect("receiver decrypts after missed messages");
638 assert_eq!(plaintext, b"payload");
639
640 let mut replacement = Device::new(sender_backend.clone());
641 let same_process = group_encrypt(&mut replacement, &name, b"same-process", &mut rng)
642 .await
643 .expect("encrypt after same-process replacement");
644 assert_eq!(same_process.iteration(), 33);
645
646 let durable = sender_backend
647 .get_sender_key(name.cache_key())
648 .await
649 .expect("read durable sender key")
650 .expect("durable sender key");
651 let recovered = SenderKeyRecord::deserialize(&durable).expect("recovery reload");
652 assert_eq!(
653 recovered
654 .sender_key_state()
655 .expect("sender-key state")
656 .sender_chain_key()
657 .expect("sender chain")
658 .iteration(),
659 wacore::libsignal::protocol::consts::SENDER_CHAIN_RESERVATION_BATCH
660 );
661 }
662
663 #[tokio::test]
666 async fn load_signed_prekey_falls_back_to_backend_for_rotated_out_id() {
667 use buffa::Message;
668
669 let backend = crate::test_utils::create_test_backend().await;
670 let device = Device::new(backend.clone());
671
672 let current = device.signed_pre_key_id;
673 let old_id = current + 7; let kp = wacore::libsignal::protocol::KeyPair::generate(&mut rand::make_rng::<
676 rand::rngs::StdRng,
677 >());
678 let record = record_helpers::new_signed_pre_key_record(
679 old_id,
680 &kp,
681 [9u8; 64],
682 wacore::time::now_utc(),
683 );
684 backend
685 .store_signed_prekey(old_id, &record.encode_to_vec())
686 .await
687 .expect("store retained signed pre-key");
688
689 let loaded = SignedPreKeyStore::load_signed_prekey(&device, old_id)
690 .await
691 .expect("load must not error")
692 .expect("rotated-out key must load from backend");
693 assert_eq!(loaded.id, Some(old_id));
694 assert_eq!(
695 loaded.public_key.as_deref(),
696 Some(kp.public_key.public_key_bytes())
697 );
698
699 let missing = SignedPreKeyStore::load_signed_prekey(&device, current + 999)
701 .await
702 .expect("load must not error");
703 assert!(missing.is_none());
704
705 assert!(
708 SignedPreKeyStore::contains_signed_prekey(&device, current)
709 .await
710 .expect("contains current")
711 );
712 assert!(
713 SignedPreKeyStore::contains_signed_prekey(&device, old_id)
714 .await
715 .expect("contains rotated-out")
716 );
717 assert!(
718 !SignedPreKeyStore::contains_signed_prekey(&device, current + 999)
719 .await
720 .expect("contains unknown")
721 );
722 }
723}