1use super::*;
4use anyhow::Context;
5
6#[cfg_attr(
10 feature = "tracing",
11 tracing::instrument(name = "wa.send.encrypt_group", level = "debug", skip_all, err(Debug))
12)]
13pub async fn encrypt_group_message<S, R>(
14 sender_key_store: &mut S,
15 sender_key_name: &SenderKeyName,
16 plaintext: &[u8],
17 csprng: &mut R,
18) -> Result<SenderKeyMessage>
19where
20 S: SenderKeyStore + ?Sized,
21 R: Rng + CryptoRng,
22{
23 crate::libsignal::protocol::group_encrypt(sender_key_store, sender_key_name, plaintext, csprng)
29 .await
30 .context("group encrypt failed")
31}
32
33pub trait CloneableSessionStore: crate::libsignal::protocol::SessionStore {
40 fn clone_box(&self) -> Box<dyn CloneableSessionStore + Send + Sync>;
41}
42
43impl<T> CloneableSessionStore for T
44where
45 T: crate::libsignal::protocol::SessionStore + Clone + Send + Sync + 'static,
46{
47 fn clone_box(&self) -> Box<dyn CloneableSessionStore + Send + Sync> {
48 Box::new(self.clone())
49 }
50}
51
52pub trait CloneableIdentityStore: crate::libsignal::protocol::IdentityKeyStore {
54 fn clone_box(&self) -> Box<dyn CloneableIdentityStore + Send + Sync>;
55}
56
57impl<T> CloneableIdentityStore for T
58where
59 T: crate::libsignal::protocol::IdentityKeyStore + Clone + Send + Sync + 'static,
60{
61 fn clone_box(&self) -> Box<dyn CloneableIdentityStore + Send + Sync> {
62 Box::new(self.clone())
63 }
64}
65
66pub struct SignalStores<'a> {
70 pub sender_key_store: &'a mut (dyn SenderKeyStore + Send + Sync),
71 pub session_store: &'a mut (dyn CloneableSessionStore + Send + Sync),
72 pub identity_store: &'a mut (dyn CloneableIdentityStore + Send + Sync),
73 pub prekey_store: &'a mut (dyn crate::libsignal::protocol::PreKeyStore + Send + Sync),
74 pub signed_prekey_store: &'a (dyn crate::libsignal::protocol::SignedPreKeyStore + Send + Sync),
75}
76
77pub(crate) const UNREGISTERED_DEVICE_CODE: u16 = 406;
82
83pub(crate) fn is_device_unregistered_error(err: &anyhow::Error) -> bool {
84 crate::request::ServerErrorCode::from_anyhow(err)
85 .is_some_and(|e| e.code == UNREGISTERED_DEVICE_CODE)
86}
87
88pub struct EncryptResult {
89 pub participant_nodes: Vec<Node>,
90 pub includes_prekey_message: bool,
91 pub encrypted_devices: Vec<Jid>,
92 pub had_unregistered_device: bool,
94 pub rejected_devices: Vec<Jid>,
97}
98
99pub(crate) struct EncryptAttempt {
100 pub result: EncryptResult,
101 pub first_error: Option<anyhow::Error>,
102}
103
104pub struct EncryptedDevice {
107 pub device_jid: Jid,
108 pub enc_type: &'static str,
110 pub is_prekey: bool,
111 pub ciphertext: Vec<u8>,
112}
113
114pub struct EncryptForDevicesRaw {
117 pub devices: Vec<EncryptedDevice>,
118 pub includes_prekey_message: bool,
119 pub had_unregistered_device: bool,
121 pub rejected_devices: Vec<Jid>,
123}
124
125struct RawEncryptAttempt {
126 result: EncryptForDevicesRaw,
127 first_error: Option<anyhow::Error>,
128}
129
130pub fn needs_device_identity(
138 includes_prekey: bool,
139 account: Option<&wa::ADVSignedDeviceIdentity>,
140) -> Result<Option<Vec<u8>>> {
141 if !includes_prekey {
142 return Ok(None);
143 }
144 let acc = account
145 .ok_or_else(|| anyhow!("pkmsg requires <device-identity> but no ADV account is present"))?;
146 Ok(Some(waproto::codec::adv_signed_device_identity_to_vec(acc)))
147}
148
149const ENCRYPT_FANOUT_CONCURRENCY: usize = 16;
153
154struct EncryptOneResult {
156 enc_type: &'static str,
157 is_prekey: bool,
158 ciphertext: Vec<u8>,
159}
160
161#[derive(Debug, thiserror::Error)]
166#[error("spawned task did not produce a result (panic or runtime shutdown)")]
167struct SpawnCanceled;
168
169struct Spawned<T> {
174 rx: futures::channel::oneshot::Receiver<T>,
175 abort: Option<AbortHandle>,
176}
177
178impl<T> Future for Spawned<T> {
179 type Output = std::result::Result<T, SpawnCanceled>;
180
181 fn poll(
182 mut self: std::pin::Pin<&mut Self>,
183 cx: &mut std::task::Context<'_>,
184 ) -> std::task::Poll<Self::Output> {
185 match std::pin::Pin::new(&mut self.rx).poll(cx) {
186 std::task::Poll::Ready(Ok(value)) => {
187 if let Some(handle) = self.abort.take() {
190 handle.detach();
191 }
192 std::task::Poll::Ready(Ok(value))
193 }
194 std::task::Poll::Ready(Err(_)) => {
195 if let Some(handle) = self.abort.take() {
196 handle.detach();
197 }
198 std::task::Poll::Ready(Err(SpawnCanceled))
199 }
200 std::task::Poll::Pending => std::task::Poll::Pending,
201 }
202 }
203}
204
205impl<T> Drop for Spawned<T> {
206 fn drop(&mut self) {
207 if let Some(handle) = self.abort.take() {
211 handle.abort();
212 }
213 }
214}
215
216#[cfg(not(target_arch = "wasm32"))]
221fn spawn_oneshot<F, T>(
222 rt: &dyn Runtime,
223 fut: F,
224) -> impl Future<Output = std::result::Result<T, SpawnCanceled>> + Send + 'static
225where
226 F: Future<Output = T> + Send + 'static,
227 T: Send + 'static,
228{
229 let (tx, rx) = futures::channel::oneshot::channel();
230 let abort = rt.spawn(Box::pin(async move {
231 let _ = tx.send(fut.await);
232 }));
233 Spawned {
234 rx,
235 abort: Some(abort),
236 }
237}
238
239#[cfg(target_arch = "wasm32")]
240fn spawn_oneshot<F, T>(
241 rt: &dyn Runtime,
242 fut: F,
243) -> impl Future<Output = std::result::Result<T, SpawnCanceled>> + 'static
244where
245 F: Future<Output = T> + 'static,
246 T: 'static,
247{
248 let (tx, rx) = futures::channel::oneshot::channel();
249 let abort = rt.spawn(Box::pin(async move {
250 let _ = tx.send(fut.await);
251 }));
252 Spawned {
253 rx,
254 abort: Some(abort),
255 }
256}
257
258async fn encrypt_one_device(
263 plaintext: &[u8],
264 addr: &ProtocolAddress,
265 session_store: &mut dyn crate::libsignal::protocol::SessionStore,
266 identity_store: &mut dyn crate::libsignal::protocol::IdentityKeyStore,
267 device_jid: Jid,
268) -> (Jid, Result<Option<EncryptOneResult>>) {
269 match message_encrypt(plaintext, addr, session_store, identity_store).await {
270 Ok(encrypted_payload) => {
271 let Some((enc_type, is_prekey, serialized_bytes)) =
272 extract_ciphertext(encrypted_payload)
273 else {
274 return (device_jid, Ok(None));
275 };
276 (
277 device_jid,
278 Ok(Some(EncryptOneResult {
279 enc_type,
280 is_prekey,
281 ciphertext: serialized_bytes.into(),
283 })),
284 )
285 }
286 Err(error) => (
287 device_jid,
288 Err(anyhow::Error::new(error).context(format!("failed to encrypt for {addr}"))),
289 ),
290 }
291}
292
293fn push_raw_result(
297 (device_jid, res): (Jid, Result<Option<EncryptOneResult>>),
298 devices: &mut Vec<EncryptedDevice>,
299 includes_prekey_message: &mut bool,
300 first_error: &mut Option<anyhow::Error>,
301) {
302 match res {
303 Ok(Some(one)) => {
304 *includes_prekey_message |= one.is_prekey;
305 devices.push(EncryptedDevice {
306 device_jid,
307 enc_type: one.enc_type,
308 is_prekey: one.is_prekey,
309 ciphertext: one.ciphertext,
310 });
311 }
312 Ok(None) => {}
313 Err(error) => {
314 log::warn!("Failed to encrypt for device: {error:#}. Skipping.");
315 if first_error.is_none() {
316 *first_error = Some(error);
317 }
318 }
319 }
320}
321
322fn encrypted_device_to_participant_node(
326 one: EncryptedDevice,
327 mediatype: Option<&str>,
328 hide_decrypt_fail: bool,
329) -> Node {
330 let mut enc_builder = NodeBuilder::new("enc")
331 .attr("v", stanza::ENC_VERSION)
332 .attr("type", one.enc_type);
333 if let Some(mt) = mediatype {
336 enc_builder = enc_builder.attr("mediatype", mt);
337 }
338 if hide_decrypt_fail {
339 enc_builder = enc_builder.attr("decrypt-fail", "hide");
340 }
341 let enc_node = enc_builder.bytes(one.ciphertext).build();
342 NodeBuilder::new("to")
343 .attr("jid", one.device_jid)
344 .children([enc_node])
345 .build()
346}
347
348pub async fn encrypt_for_devices(
362 runtime: &dyn Runtime,
363 stores: &mut SignalStores<'_>,
364 resolver: &dyn SendContextResolver,
365 devices: &[Jid],
366 plaintext_to_encrypt: &[u8],
367 hide_decrypt_fail: bool,
368 mediatype: Option<&str>,
369) -> Result<EncryptResult> {
370 let plan = ensure_sessions_for_devices(runtime, stores, resolver, devices).await?;
371 encrypt_for_devices_with_sessions(
372 runtime,
373 stores,
374 devices,
375 plaintext_to_encrypt,
376 hide_decrypt_fail,
377 mediatype,
378 plan,
379 )
380 .await
381}
382
383pub struct EncryptFanoutSummary {
386 pub includes_prekey_message: bool,
387 pub had_unregistered_device: bool,
389}
390
391#[allow(clippy::too_many_arguments)]
401pub async fn encrypt_for_devices_into(
402 runtime: &dyn Runtime,
403 stores: &mut SignalStores<'_>,
404 resolver: &dyn SendContextResolver,
405 devices: &[Jid],
406 plaintext_to_encrypt: &[u8],
407 hide_decrypt_fail: bool,
408 mediatype: Option<&str>,
409 participant_nodes: &mut Vec<Node>,
410) -> Result<EncryptFanoutSummary> {
411 let plan = ensure_sessions_for_devices(runtime, stores, resolver, devices).await?;
412 let RawEncryptAttempt { result: raw, .. } = encrypt_for_devices_with_sessions_raw_detailed(
415 runtime,
416 stores,
417 devices,
418 plaintext_to_encrypt,
419 plan,
420 )
421 .await?;
422
423 participant_nodes.reserve(raw.devices.len());
424 for one in raw.devices {
425 participant_nodes.push(encrypted_device_to_participant_node(
426 one,
427 mediatype,
428 hide_decrypt_fail,
429 ));
430 }
431
432 Ok(EncryptFanoutSummary {
433 includes_prekey_message: raw.includes_prekey_message,
434 had_unregistered_device: raw.had_unregistered_device,
435 })
436}
437
438pub struct SessionPlan {
444 device_count: usize,
448 encryption_overrides: Vec<Option<Jid>>,
451 pub had_unregistered_device: bool,
452 pub rejected_devices: Vec<Jid>,
462 first_error: Option<anyhow::Error>,
463}
464
465impl SessionPlan {
466 pub fn assume_ready(device_count: usize) -> Self {
471 Self {
472 device_count,
473 encryption_overrides: Vec::new(),
474 had_unregistered_device: false,
475 rejected_devices: Vec::new(),
476 first_error: None,
477 }
478 }
479}
480
481fn encryption_override_at(overrides: &[Option<Jid>], index: usize) -> Option<&Jid> {
485 overrides.get(index).and_then(Option::as_ref)
486}
487
488fn record_encryption_override(
494 overrides: &mut Vec<Option<Jid>>,
495 device_count: usize,
496 index: usize,
497 jid: Jid,
498) {
499 if overrides.is_empty() {
500 overrides.resize(device_count, None);
501 }
502 overrides[index] = Some(jid);
503}
504
505#[cfg_attr(feature = "tracing", tracing::instrument(name = "wa.send.ensure_sessions", level = "debug", skip_all, fields(count = devices.len()), err(Debug)))]
510pub async fn ensure_sessions_for_devices(
511 runtime: &dyn Runtime,
512 stores: &mut SignalStores<'_>,
513 resolver: &dyn SendContextResolver,
514 devices: &[Jid],
515) -> Result<SessionPlan> {
516 let mut encryption_overrides: Vec<Option<Jid>> = Vec::new();
524 let mut indices_needing_prekeys: Vec<usize> = Vec::new();
526 let mut had_406 = false;
527 let mut rejected_devices: Vec<Jid> = Vec::new();
528 let mut first_error = None;
529
530 let mut reusable_addr = crate::types::jid::make_reusable_protocol_address();
531
532 for (idx, device_jid) in devices.iter().enumerate() {
533 if device_jid.is_pn()
537 && let Some(lid_user) = resolver.get_lid_for_phone(&device_jid.user).await
538 {
539 let lid_jid = Jid::lid_device(lid_user, device_jid.device);
541 lid_jid.reset_protocol_address(&mut reusable_addr);
542
543 if wacore_libsignal::protocol::has_session(stores.session_store, &reusable_addr).await?
544 {
545 log::debug!(
546 "Using LID session {} for PN {} (LID-first lookup)",
547 lid_jid.observe(),
548 device_jid.observe()
549 );
550 record_encryption_override(&mut encryption_overrides, devices.len(), idx, lid_jid);
551 continue;
552 }
553 }
554
555 device_jid.reset_protocol_address(&mut reusable_addr);
556 if wacore_libsignal::protocol::has_session(stores.session_store, &reusable_addr).await? {
557 continue;
558 }
559
560 if device_jid.is_pn()
564 && let Some(lid_user) = resolver.get_lid_for_phone(&device_jid.user).await
565 {
566 let lid_jid = Jid::lid_device(lid_user, device_jid.device);
567 log::debug!(
568 "Will create LID session {} for PN {} (no existing session)",
569 lid_jid.observe(),
570 device_jid.observe()
571 );
572 record_encryption_override(&mut encryption_overrides, devices.len(), idx, lid_jid);
573 }
574 indices_needing_prekeys.push(idx);
575 }
576
577 if !indices_needing_prekeys.is_empty() {
578 log::debug!(
579 "Fetching prekeys for {} devices without sessions",
580 indices_needing_prekeys.len()
581 );
582 let jids_for_fetch: Vec<Jid> = indices_needing_prekeys
586 .iter()
587 .map(|&i| devices[i].clone())
588 .collect();
589 let prekey_bundles = match resolver
598 .fetch_prekeys_for_identity_check(&jids_for_fetch)
599 .await
600 {
601 Ok(outcome) => {
602 rejected_devices.extend(
603 outcome
604 .rejected
605 .iter()
606 .filter(|device| device.code == UNREGISTERED_DEVICE_CODE)
607 .map(|device| device.jid.clone()),
608 );
609 if !rejected_devices.is_empty() {
610 log::debug!(
611 "prekey fetch rejected {} of {} device(s) by name",
612 rejected_devices.len(),
613 jids_for_fetch.len()
614 );
615 had_406 = true;
616 }
617 outcome.bundles
618 }
619 Err(e) if is_device_unregistered_error(&e) => {
620 log::debug!(
623 "Prekey fetch returned 406 for {} device(s); skipping them this round",
624 jids_for_fetch.len()
625 );
626 had_406 = true;
627 first_error = Some(e);
631 std::collections::HashMap::new()
632 }
633 Err(e) => return Err(e),
634 };
635
636 let prekey_bundles = std::sync::Arc::new(prekey_bundles);
644 let total = indices_needing_prekeys.len();
645 let mut next_spawn = 0usize;
646
647 let make_session_task = |spawn_idx: usize| {
648 let idx = indices_needing_prekeys[spawn_idx];
649 let lookup_jid = devices[idx].clone();
650 let encryption_jid = encryption_override_at(&encryption_overrides, idx)
651 .cloned()
652 .unwrap_or_else(|| lookup_jid.clone());
653
654 let bundles = prekey_bundles.clone();
655 let mut session_store = stores.session_store.clone_box();
656 let mut identity_store = stores.identity_store.clone_box();
657
658 spawn_oneshot(runtime, async move {
659 let mut addr = crate::types::jid::make_reusable_protocol_address();
660 encryption_jid.reset_protocol_address(&mut addr);
661
662 let Some(bundle) = bundles.get(&lookup_jid) else {
663 log::debug!(
666 "No pre-key bundle returned for device {}. This device will be skipped for encryption.",
667 addr
668 );
669 return Ok::<Option<Jid>, anyhow::Error>(None);
670 };
671
672 let mut rng = rand::make_rng::<rand::rngs::StdRng>();
673 match process_prekey_bundle(
677 &addr,
678 &mut *session_store,
679 &mut *identity_store,
680 bundle,
681 &mut rng,
682 UsePQRatchet::No,
683 )
684 .await
685 {
686 Ok(IdentityChange::ReplacedExisting) => Ok(Some(encryption_jid)),
689 Ok(IdentityChange::NewOrUnchanged) => Ok(None),
690 Err(error) => Err(anyhow::Error::new(error)
691 .context(format!("failed to process pre-key bundle for {addr}"))),
692 }
693 })
694 };
695
696 let mut in_flight: FuturesUnordered<_> = FuturesUnordered::new();
697 while next_spawn < total && in_flight.len() < ENCRYPT_FANOUT_CONCURRENCY {
698 in_flight.push(make_session_task(next_spawn));
699 next_spawn += 1;
700 }
701 while let Some(spawn_result) = in_flight.next().await {
702 match spawn_result {
703 Ok(Ok(Some(changed_jid))) => resolver.on_local_identity_change(&changed_jid),
706 Ok(Ok(None)) => {}
707 Ok(Err(e)) => {
712 log::warn!("Group session setup failed for a device, skipping it: {e}");
713 if first_error.is_none() {
714 first_error = Some(e);
715 }
716 }
717 Err(error) => {
718 log::warn!(
719 "Session-establishment task did not deliver a result; skipping device."
720 );
721 if first_error.is_none() {
722 first_error = Some(anyhow::Error::new(error));
723 }
724 }
725 }
726 if next_spawn < total {
727 in_flight.push(make_session_task(next_spawn));
728 next_spawn += 1;
729 }
730 }
731 }
732
733 Ok(SessionPlan {
734 device_count: devices.len(),
735 encryption_overrides,
736 had_unregistered_device: had_406,
737 rejected_devices,
738 first_error,
739 })
740}
741
742pub async fn encrypt_for_devices_with_sessions(
749 runtime: &dyn Runtime,
750 stores: &mut SignalStores<'_>,
751 devices: &[Jid],
752 plaintext_to_encrypt: &[u8],
753 hide_decrypt_fail: bool,
754 mediatype: Option<&str>,
755 plan: SessionPlan,
756) -> Result<EncryptResult> {
757 Ok(encrypt_for_devices_with_sessions_detailed(
758 runtime,
759 stores,
760 devices,
761 plaintext_to_encrypt,
762 hide_decrypt_fail,
763 mediatype,
764 plan,
765 )
766 .await?
767 .result)
768}
769
770#[cfg_attr(feature = "tracing", tracing::instrument(name = "wa.send.encrypt_fanout", level = "debug", skip_all, fields(count = devices.len()), err(Debug)))]
771pub(crate) async fn encrypt_for_devices_with_sessions_detailed(
772 runtime: &dyn Runtime,
773 stores: &mut SignalStores<'_>,
774 devices: &[Jid],
775 plaintext_to_encrypt: &[u8],
776 hide_decrypt_fail: bool,
777 mediatype: Option<&str>,
778 plan: SessionPlan,
779) -> Result<EncryptAttempt> {
780 let RawEncryptAttempt {
781 result: raw,
782 first_error,
783 } = encrypt_for_devices_with_sessions_raw_detailed(
784 runtime,
785 stores,
786 devices,
787 plaintext_to_encrypt,
788 plan,
789 )
790 .await?;
791
792 let mut participant_nodes = Vec::with_capacity(raw.devices.len());
796 let mut encrypted_devices = Vec::with_capacity(raw.devices.len());
797 for one in raw.devices {
798 encrypted_devices.push(one.device_jid.clone());
799 participant_nodes.push(encrypted_device_to_participant_node(
800 one,
801 mediatype,
802 hide_decrypt_fail,
803 ));
804 }
805
806 Ok(EncryptAttempt {
807 result: EncryptResult {
808 participant_nodes,
809 includes_prekey_message: raw.includes_prekey_message,
810 encrypted_devices,
811 had_unregistered_device: raw.had_unregistered_device,
812 rejected_devices: raw.rejected_devices.clone(),
813 },
814 first_error,
815 })
816}
817
818pub async fn encrypt_for_devices_with_sessions_raw(
826 runtime: &dyn Runtime,
827 stores: &mut SignalStores<'_>,
828 devices: &[Jid],
829 plaintext_to_encrypt: &[u8],
830 plan: SessionPlan,
831) -> Result<EncryptForDevicesRaw> {
832 Ok(encrypt_for_devices_with_sessions_raw_detailed(
833 runtime,
834 stores,
835 devices,
836 plaintext_to_encrypt,
837 plan,
838 )
839 .await?
840 .result)
841}
842
843#[cfg_attr(feature = "tracing", tracing::instrument(name = "wa.send.encrypt_fanout_raw", level = "debug", skip_all, fields(count = devices.len()), err(Debug)))]
844async fn encrypt_for_devices_with_sessions_raw_detailed(
845 runtime: &dyn Runtime,
846 stores: &mut SignalStores<'_>,
847 devices: &[Jid],
848 plaintext_to_encrypt: &[u8],
849 plan: SessionPlan,
850) -> Result<RawEncryptAttempt> {
851 debug_assert_eq!(
852 plan.device_count,
853 devices.len(),
854 "SessionPlan built for a different device list"
855 );
856 let SessionPlan {
857 device_count: _,
858 encryption_overrides,
859 had_unregistered_device,
860 rejected_devices,
861 mut first_error,
862 } = plan;
863
864 let mut encrypted = Vec::with_capacity(devices.len());
865 let mut includes_prekey_message = false;
866
867 if devices.len() == 1 {
871 let device_jid = devices[0].clone();
876 let addr = encryption_override_at(&encryption_overrides, 0)
877 .unwrap_or(&devices[0])
878 .to_protocol_address();
879 let res = encrypt_one_device(
880 plaintext_to_encrypt,
881 &addr,
882 &mut *stores.session_store,
883 &mut *stores.identity_store,
884 device_jid,
885 )
886 .await;
887 push_raw_result(
888 res,
889 &mut encrypted,
890 &mut includes_prekey_message,
891 &mut first_error,
892 );
893 } else {
894 let plaintext_arc: std::sync::Arc<[u8]> = std::sync::Arc::from(plaintext_to_encrypt);
899
900 let total = devices.len();
901 let num_chunks = ENCRYPT_FANOUT_CONCURRENCY.min(total);
902
903 let mut in_flight: FuturesUnordered<_> = FuturesUnordered::new();
904 for chunk_idx in 0..num_chunks {
907 let chunk_start = chunk_idx * total / num_chunks;
908 let chunk_end = (chunk_idx + 1) * total / num_chunks;
909 let jobs: Vec<(ProtocolAddress, Jid)> = (chunk_start..chunk_end)
911 .map(|idx| {
912 let addr = encryption_override_at(&encryption_overrides, idx)
913 .unwrap_or(&devices[idx])
914 .to_protocol_address();
915 (addr, devices[idx].clone())
916 })
917 .collect();
918 let plaintext = plaintext_arc.clone();
919 let mut session_store = stores.session_store.clone_box();
922 let mut identity_store = stores.identity_store.clone_box();
923
924 in_flight.push(spawn_oneshot(runtime, async move {
925 let mut out = Vec::with_capacity(jobs.len());
926 for (addr, device_jid) in jobs {
927 out.push(
928 encrypt_one_device(
929 &plaintext,
930 &addr,
931 &mut *session_store,
932 &mut *identity_store,
933 device_jid,
934 )
935 .await,
936 );
937 }
938 out
939 }));
940 }
941 while let Some(spawn_result) = in_flight.next().await {
942 match spawn_result {
943 Ok(results) => {
944 for res in results {
945 push_raw_result(
946 res,
947 &mut encrypted,
948 &mut includes_prekey_message,
949 &mut first_error,
950 );
951 }
952 }
953 Err(error) => {
954 log::warn!(
957 "Encrypt chunk did not deliver a result; up to ~{} device(s) skipped this send.",
958 total.div_ceil(num_chunks)
959 );
960 if first_error.is_none() {
961 first_error = Some(anyhow::Error::new(error));
962 }
963 }
964 }
965 }
966 }
967
968 Ok(RawEncryptAttempt {
969 result: EncryptForDevicesRaw {
970 devices: encrypted,
971 includes_prekey_message,
972 had_unregistered_device,
973 rejected_devices,
974 },
975 first_error,
976 })
977}
978
979#[cfg(test)]
980mod encryption_override_tests {
981 use super::{SessionPlan, encryption_override_at, record_encryption_override};
982 use wacore_binary::Jid;
983
984 fn lid(user: &str, device: u16) -> Jid {
985 Jid::lid_device(user.to_owned(), device)
986 }
987
988 #[test]
991 fn an_empty_map_answers_every_index_without_allocating() {
992 let overrides: Vec<Option<Jid>> = Vec::new();
993 assert_eq!(overrides.capacity(), 0, "no override must mean no buffer");
994 for index in [0, 1, 7, usize::MAX] {
995 assert!(encryption_override_at(&overrides, index).is_none());
996 }
997
998 let plan = SessionPlan::assume_ready(4);
999 assert!(
1000 plan.encryption_overrides.is_empty(),
1001 "a plan that overrides nothing must carry no override buffer"
1002 );
1003 assert_eq!(plan.device_count, 4, "the slice length is still recorded");
1004 }
1005
1006 #[test]
1009 fn recording_materializes_the_whole_map_once() {
1010 let mut overrides: Vec<Option<Jid>> = Vec::new();
1011 record_encryption_override(&mut overrides, 3, 2, lid("100000000000001", 5));
1012 assert_eq!(overrides.len(), 3);
1013 assert!(encryption_override_at(&overrides, 0).is_none());
1014 assert!(encryption_override_at(&overrides, 1).is_none());
1015 assert_eq!(
1016 encryption_override_at(&overrides, 2),
1017 Some(&lid("100000000000001", 5))
1018 );
1019
1020 record_encryption_override(&mut overrides, 3, 0, lid("100000000000002", 0));
1022 assert_eq!(overrides.len(), 3);
1023 assert_eq!(
1024 encryption_override_at(&overrides, 0),
1025 Some(&lid("100000000000002", 0))
1026 );
1027 assert_eq!(
1028 encryption_override_at(&overrides, 2),
1029 Some(&lid("100000000000001", 5))
1030 );
1031
1032 record_encryption_override(&mut overrides, 3, 2, lid("100000000000003", 1));
1034 assert_eq!(overrides.len(), 3);
1035 assert_eq!(
1036 encryption_override_at(&overrides, 2),
1037 Some(&lid("100000000000003", 1))
1038 );
1039 }
1040
1041 #[test]
1044 fn a_single_device_plan_records_and_reads_index_zero() {
1045 let mut overrides: Vec<Option<Jid>> = Vec::new();
1046 assert!(encryption_override_at(&overrides, 0).is_none());
1047 record_encryption_override(&mut overrides, 1, 0, lid("100000000000009", 33));
1048 assert_eq!(
1049 encryption_override_at(&overrides, 0),
1050 Some(&lid("100000000000009", 33))
1051 );
1052 assert!(encryption_override_at(&overrides, 1).is_none());
1053 }
1054}