1use std::{
19 collections::BTreeMap,
20 future::Future,
21 sync::{Arc, Mutex, atomic::Ordering},
22};
23
24use matrix_sdk_base::crypto::types::events::room::encrypted::EncryptedToDeviceEvent;
25use matrix_sdk_test::test_json;
26use ruma::{
27 CrossSigningKeyId, DeviceId, MilliSecondsSinceUnixEpoch, OneTimeKeyAlgorithm, OwnedDeviceId,
28 OwnedOneTimeKeyId, OwnedUserId, UserId,
29 api::client::{
30 keys::upload_signatures::v3::SignedKeys, to_device::send_event_to_device::v3::Messages,
31 },
32 encryption::{CrossSigningKey, DeviceKeys, OneTimeKey},
33 events::AnyToDeviceEvent,
34 owned_device_id, owned_user_id,
35 serde::Raw,
36 to_device::DeviceIdOrAllDevices,
37};
38use serde_json::json;
39use strass::assert_let;
40use tracing::Instrument;
41use wiremock::{
42 Mock, MockGuard, Request, ResponseTemplate,
43 matchers::{method, path_regex},
44};
45
46use crate::{
47 Client,
48 test_utils::{
49 client::MockClientBuilder,
50 mocks::{Keys, MatrixMockServer},
51 },
52};
53
54pub type PendingToDeviceMessages =
57 BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, Vec<Raw<AnyToDeviceEvent>>>>;
58
59impl MatrixMockServer {
82 pub fn client_builder_for_crypto_end_to_end(
86 &self,
87 user_id: &UserId,
88 device_id: &DeviceId,
89 ) -> MockClientBuilder {
90 let next = self.token_counter.fetch_add(1, Ordering::Relaxed);
92 let access_token = format!("TOKEN_{next}");
93
94 {
95 let mut mappings = self.token_to_user_id_map.lock().unwrap();
96 let auth_string = format!("Bearer {access_token}");
97 mappings.insert(auth_string, user_id.to_owned());
98 }
99
100 MockClientBuilder::new(Some(&self.server.uri())).logged_in_with_token(
101 access_token,
102 user_id.to_owned(),
103 device_id.to_owned(),
104 )
105 }
106
107 pub fn exhaust_one_time_keys(&self, user_id: OwnedUserId, device_id: OwnedDeviceId) {
109 let mut keys = self.keys.lock().unwrap();
110 let known_otks = &mut keys.one_time_keys;
111 known_otks.entry(user_id).or_default().entry(device_id).or_default().clear();
112 }
113
114 pub async fn exchange_e2ee_identities(&self, alice: &Client, bob: &Client) {
117 let alice_user_id = alice.user_id().expect("Alice should have a user ID configured");
118 let bob_user_id = bob.user_id().expect("Bob should have a user ID configured");
119
120 let alice_span = tracing::info_span!("alice", user_id=%alice_user_id);
121 let bob_span = tracing::info_span!("bob", user_id=%bob_user_id);
122
123 alice.update_tracked_users_for_testing([bob_user_id]).instrument(alice_span.clone()).await;
125
126 bob.update_tracked_users_for_testing([alice_user_id]).instrument(bob_span.clone()).await;
130
131 self.mock_sync().ok_and_run(alice, |_x| {}).instrument(alice_span.clone()).await;
133 self.mock_sync().ok_and_run(bob, |_x| {}).instrument(bob_span).await;
134
135 self.mock_sync().ok_and_run(alice, |_x| {}).instrument(alice_span).await;
138 }
139
140 pub async fn set_up_alice_and_bob_for_encryption(&self) -> (Client, Client) {
143 let alice_user_id = owned_user_id!("@alice:example.org");
144 let alice_device_id = owned_device_id!("4L1C3");
145
146 let alice = self
147 .client_builder_for_crypto_end_to_end(&alice_user_id, &alice_device_id)
148 .build()
149 .await;
150
151 let bob_user_id = owned_user_id!("@bob:example.org");
152 let bob_device_id = owned_device_id!("B0B0B0B0B");
153 let bob =
154 self.client_builder_for_crypto_end_to_end(&bob_user_id, &bob_device_id).build().await;
155
156 self.exchange_e2ee_identities(&alice, &bob).await;
157
158 (alice, bob)
159 }
160
161 pub async fn set_up_carl_for_encryption(&self, alice: &Client, bob: &Client) -> Client {
163 let carl_user_id = owned_user_id!("@carlg:example.org");
164 let carl_device_id = owned_device_id!("CARL_DEVICE");
165
166 let carl =
167 self.client_builder_for_crypto_end_to_end(&carl_user_id, &carl_device_id).build().await;
168
169 self.mock_sync().ok_and_run(&carl, |_| {}).await;
171
172 alice.update_tracked_users_for_testing([carl.user_id().unwrap()]).await;
174
175 bob.update_tracked_users_for_testing([carl.user_id().unwrap()]).await;
177
178 {
181 self.mock_sync().ok_and_run(alice, |_| {}).await;
182 self.mock_sync().ok_and_run(bob, |_| {}).await;
183 }
184
185 carl.update_tracked_users_for_testing([alice.user_id().unwrap(), bob.user_id().unwrap()])
187 .await;
188
189 self.mock_sync().ok_and_run(alice, |_| {}).await;
191
192 carl
193 }
194
195 pub async fn set_up_new_device_for_encryption(
211 &self,
212 existing_client: &Client,
213 device_id: &DeviceId,
214 clients_to_update: Vec<&Client>,
215 ) -> Client {
216 let user_id = existing_client.user_id().unwrap().to_owned();
217 let new_device_id = device_id.to_owned();
218
219 let new_client =
220 self.client_builder_for_crypto_end_to_end(&user_id, &new_device_id).build().await;
221
222 self.mock_sync().ok_and_run(&new_client, |_| {}).await;
224
225 self.mock_sync()
227 .ok_and_run(existing_client, |builder| {
228 builder.add_change_device(&user_id);
229 })
230 .await;
231
232 for client_to_update in clients_to_update {
233 self.mock_sync()
234 .ok_and_run(client_to_update, |builder| {
235 builder.add_change_device(&user_id);
236 })
237 .await;
238 }
239
240 new_client
241 }
242
243 pub async fn mock_crypto_endpoints_preset(&self) {
246 let keys = &self.keys;
247 let token_map = &self.token_to_user_id_map;
248
249 Mock::given(method("POST"))
250 .and(path_regex(r"^/_matrix/client/.*/keys/query"))
251 .respond_with(mock_keys_query(keys.clone()))
252 .mount(&self.server)
253 .await;
254
255 Mock::given(method("POST"))
256 .and(path_regex(r"^/_matrix/client/.*/keys/upload"))
257 .respond_with(mock_keys_upload(keys.clone(), token_map.clone()))
258 .mount(&self.server)
259 .await;
260
261 Mock::given(method("POST"))
262 .and(path_regex(r"^/_matrix/client/.*/keys/device_signing/upload"))
263 .respond_with(mock_keys_device_signing_upload(keys.clone()))
264 .mount(&self.server)
265 .await;
266
267 Mock::given(method("POST"))
268 .and(path_regex(r"^/_matrix/client/.*/keys/signatures/upload"))
269 .respond_with(mock_keys_signature_upload(keys.clone()))
270 .mount(&self.server)
271 .await;
272
273 Mock::given(method("POST"))
274 .and(path_regex(r"^/_matrix/client/.*/keys/claim"))
275 .respond_with(mock_keys_claimed_request(keys.clone()))
276 .mount(&self.server)
277 .await;
278 }
279
280 pub async fn mock_capture_put_to_device(
346 &self,
347 sender_user_id: &UserId,
348 ) -> (MockGuard, impl Future<Output = Raw<EncryptedToDeviceEvent>> + use<>) {
349 let (tx, rx) = tokio::sync::oneshot::channel();
350 let tx = Arc::new(Mutex::new(Some(tx)));
351
352 let sender = sender_user_id.to_owned();
353 let guard = Mock::given(method("PUT"))
354 .and(path_regex(r"^/_matrix/client/.*/sendToDevice/m.room.encrypted/.*"))
355 .respond_with(move |req: &Request| {
356 #[derive(Debug, serde::Deserialize)]
357 struct Parameters {
358 messages: Messages,
359 }
360
361 let params: Parameters = req.body_json().unwrap();
362
363 let (_, device_to_content) = params.messages.first_key_value().unwrap();
364 let content = device_to_content.first_key_value().unwrap().1;
365
366 let event = json!({
367 "origin_server_ts": MilliSecondsSinceUnixEpoch::now(),
368 "sender": sender,
369 "type": "m.room.encrypted",
370 "content": content,
371 });
372 let event: Raw<EncryptedToDeviceEvent> = serde_json::from_value(event).unwrap();
373
374 if let Ok(mut guard) = tx.lock()
375 && let Some(tx) = guard.take()
376 {
377 let _ = tx.send(event);
378 }
379
380 ResponseTemplate::new(200).set_body_json(&*test_json::EMPTY)
381 })
382 .expect(1)
384 .named("send_to_device")
385 .mount_as_scoped(self.server())
386 .await;
387
388 let future =
389 async move { rx.await.expect("Failed to receive captured value - sender was dropped") };
390
391 (guard, future)
392 }
393
394 pub async fn mock_capture_put_to_device_then_sync_back<'a>(
412 &'a self,
413 sender_user_id: &UserId,
414 recipient: &'a Client,
415 ) -> impl Future<Output = Raw<EncryptedToDeviceEvent>> + 'a {
416 let (guard, sent_event) = self.mock_capture_put_to_device(sender_user_id).await;
417
418 async {
419 let sent_event = sent_event.await;
420 drop(guard);
421 self.mock_sync()
422 .ok_and_run(recipient, |sync_builder| {
423 sync_builder.add_to_device_event(sent_event.deserialize_as().unwrap());
424 })
425 .await;
426
427 sent_event
428 }
429 }
430
431 pub async fn capture_put_to_device_traffic(
435 &self,
436 sender_user_id: &UserId,
437 to_device_queue: Arc<Mutex<PendingToDeviceMessages>>,
438 ) -> MockGuard {
439 let sender = sender_user_id.to_owned();
440
441 Mock::given(method("PUT"))
442 .and(path_regex(r"^/_matrix/client/.*/sendToDevice/([^/]+)/.*"))
443 .respond_with(move |req: &Request| {
444 #[derive(Debug, serde::Deserialize)]
445 struct Parameters {
446 messages: Messages,
447 }
448
449 let params: Parameters = req.body_json().unwrap();
450 let messages = params.messages;
451
452 let event_type = req
454 .url
455 .path_segments()
456 .and_then(|segments| segments.rev().nth(1))
457 .expect("Event type should be captured in the path");
458
459 let mut to_device_queue = to_device_queue.lock().unwrap();
460 for (user_id, device_map) in messages.iter() {
461 for (device_id, content) in device_map.iter() {
462 assert_let!(DeviceIdOrAllDevices::DeviceId(device_id) = device_id);
463
464 let event = json!({
465 "origin_server_ts": MilliSecondsSinceUnixEpoch::now(),
466 "sender": sender,
467 "type": event_type.to_owned(),
468 "content": content,
469 });
470
471 to_device_queue
472 .entry(user_id.to_owned())
473 .or_default()
474 .entry(device_id.to_owned())
475 .or_default()
476 .push(serde_json::from_value(event).unwrap());
477 }
478 }
479
480 ResponseTemplate::new(200).set_body_json(&*test_json::EMPTY)
481 })
482 .mount_as_scoped(self.server())
483 .await
484 }
485
486 pub async fn sync_back_pending_to_device_messages(
492 &self,
493 to_device_queue: Arc<Mutex<PendingToDeviceMessages>>,
494 recipient: &Client,
495 ) {
496 let messages_to_sync = {
497 let to_device_queue = to_device_queue.lock().unwrap();
498 let pending_messages = to_device_queue
499 .get(&recipient.user_id().unwrap().to_owned())
500 .and_then(|treemap| treemap.get(&recipient.device_id().unwrap().to_owned()));
501
502 pending_messages.cloned().unwrap_or_default()
503 };
504
505 for message in messages_to_sync {
506 self.mock_sync()
507 .ok_and_run(recipient, |sync_builder| {
508 sync_builder.add_to_device_event(message.deserialize_as().unwrap());
509 })
510 .await;
511 }
512 }
513}
514
515fn mock_keys_query(keys: Arc<Mutex<Keys>>) -> impl Fn(&Request) -> ResponseTemplate {
520 move |req| {
521 #[derive(Debug, serde::Deserialize)]
522 struct Parameters {
523 device_keys: BTreeMap<OwnedUserId, Vec<OwnedDeviceId>>,
524 }
525
526 let params: Parameters = req.body_json().unwrap();
527
528 let keys = keys.lock().unwrap();
529 let mut device_keys = keys.device.clone();
530 if !params.device_keys.is_empty() {
531 device_keys.retain(|user, key_map| {
532 if let Some(devices) = params.device_keys.get(user) {
533 if !devices.is_empty() {
534 key_map.retain(|key_id, _json| {
535 devices.iter().any(|device_id| &device_id.to_string() == key_id)
536 });
537 }
538 true
539 } else {
540 false
541 }
542 })
543 }
544
545 let master_keys = keys.master.clone();
546 let self_signing_keys = keys.self_signing.clone();
547 let user_signing_keys = keys.user_signing.clone();
548
549 ResponseTemplate::new(200).set_body_json(json!({
550 "device_keys": device_keys,
551 "master_keys": master_keys,
552 "self_signing_keys": self_signing_keys,
553 "user_signing_keys": user_signing_keys,
554 }))
555 }
556}
557
558fn mock_keys_upload(
564 keys: Arc<Mutex<Keys>>,
565 token_to_user_id_map: Arc<Mutex<BTreeMap<String, OwnedUserId>>>,
566) -> impl Fn(&Request) -> ResponseTemplate {
567 move |req: &Request| {
568 #[derive(Debug, serde::Deserialize)]
569 struct Parameters {
570 device_keys: Option<Raw<DeviceKeys>>,
571 one_time_keys: Option<BTreeMap<OwnedOneTimeKeyId, Raw<OneTimeKey>>>,
572 }
573 let bearer_token = req
574 .headers
575 .get(http::header::AUTHORIZATION)
576 .and_then(|header| header.to_str().ok())
577 .expect("This call should be authenticated");
578
579 let params: Parameters = req.body_json().unwrap();
580
581 let tokens = token_to_user_id_map.lock().unwrap();
582 let user_id = tokens.get(bearer_token)
584 .expect("Expect this token to be known, ensure you use `MatrixKeysServer::client_builder_for_crypto_end_to_end`")
585 .to_owned();
586
587 if let Some(new_device_keys) = params.device_keys {
588 let new_device_keys = new_device_keys.deserialize().unwrap();
589
590 let key_id = new_device_keys.device_id.to_string();
591 let mut keys = keys.lock().unwrap();
593 let devices = keys.device.entry(new_device_keys.user_id.clone()).or_default();
594
595 if let Some(device_keys) = devices.get_mut(&key_id) {
598 let mut existing = device_keys.deserialize().unwrap();
599
600 for (uid, sigs) in existing.signatures.iter_mut() {
602 if let Some(new_sigs) = new_device_keys.signatures.get(uid) {
603 sigs.extend(new_sigs.clone());
604 }
605 }
606 for (uid, sigs) in new_device_keys.signatures.iter() {
607 if !existing.signatures.contains_key(uid) {
608 existing.signatures.insert(uid.clone(), sigs.clone());
609 }
610 }
611
612 *device_keys = Raw::new(&existing).unwrap();
613 } else {
614 devices.insert(key_id, Raw::new(&new_device_keys).unwrap());
615 }
616 }
617
618 let mut keys = keys.lock().unwrap();
619
620 if let Some(otks) = params.one_time_keys {
621 for (key_id, raw_otk) in otks {
625 let otk = raw_otk.deserialize().unwrap();
626 match otk {
627 OneTimeKey::SignedKey(signed_key) => {
628 let device_id = signed_key
629 .signatures
630 .first_key_value()
631 .unwrap()
632 .1
633 .keys()
634 .next()
635 .unwrap()
636 .key_name()
637 .to_owned();
638
639 keys.one_time_keys
640 .entry(user_id.clone())
641 .or_default()
642 .entry(device_id)
643 .or_default()
644 .insert(key_id, raw_otk);
645 }
646 OneTimeKey::Key(_) => {
647 }
649 _ => {}
650 }
651 }
652 }
653
654 let otk_count = keys.one_time_keys.get(&user_id).map(|m| m.len()).unwrap_or(0);
655 ResponseTemplate::new(200).set_body_json(json!({
656 "one_time_key_counts": {
657 "signed_curve25519": otk_count,
658 }
659 }))
660 }
661}
662
663fn mock_keys_device_signing_upload(
671 keys: Arc<Mutex<Keys>>,
672) -> impl Fn(&Request) -> ResponseTemplate {
673 move |req: &Request| {
674 #[derive(Debug, serde::Deserialize)]
676 struct Parameters {
677 master_key: Option<Raw<CrossSigningKey>>,
678 self_signing_key: Option<Raw<CrossSigningKey>>,
679 user_signing_key: Option<Raw<CrossSigningKey>>,
680 }
681
682 let params: Parameters = req.body_json().unwrap();
683 assert!(params.master_key.is_some());
684 assert!(params.self_signing_key.is_some());
685 assert!(params.user_signing_key.is_some());
686
687 let mut keys = keys.lock().unwrap();
688
689 if let Some(key) = params.master_key {
690 let deserialized = key.deserialize().unwrap();
691 let user_id = deserialized.user_id;
692 keys.master.insert(user_id, key);
693 }
694
695 if let Some(key) = params.self_signing_key {
696 let deserialized = key.deserialize().unwrap();
697 let user_id = deserialized.user_id;
698 keys.self_signing.insert(user_id, key);
699 }
700
701 if let Some(key) = params.user_signing_key {
702 let deserialized = key.deserialize().unwrap();
703 let user_id = deserialized.user_id;
704 keys.user_signing.insert(user_id, key);
705 }
706
707 ResponseTemplate::new(200).set_body_json(json!({}))
708 }
709}
710
711fn mock_keys_signature_upload(keys: Arc<Mutex<Keys>>) -> impl Fn(&Request) -> ResponseTemplate {
715 move |req: &Request| {
716 #[derive(Debug, serde::Deserialize)]
717 #[serde(transparent)]
718 struct Parameters(BTreeMap<OwnedUserId, SignedKeys>);
719
720 let params: Parameters = req.body_json().unwrap();
721
722 let mut keys = keys.lock().unwrap();
723
724 for (user, signed_keys) in params.0 {
725 for (key_id, raw_key) in signed_keys.iter() {
726 if let Some(existing_master_key) = keys.master.get_mut(&user) {
728 let mut existing = existing_master_key.deserialize().unwrap();
729
730 let target = CrossSigningKeyId::from_parts(
731 ruma::SigningKeyAlgorithm::Ed25519,
732 key_id.try_into().unwrap(),
733 );
734
735 if existing.keys.contains_key(&target) {
736 let param: CrossSigningKey = serde_json::from_str(raw_key.get()).unwrap();
737
738 for (uid, sigs) in existing.signatures.iter_mut() {
739 if let Some(new_sigs) = param.signatures.get(uid) {
740 sigs.extend(new_sigs.clone());
741 }
742 }
743 for (uid, sigs) in param.signatures.iter() {
744 if !existing.signatures.contains_key(uid) {
745 existing.signatures.insert(uid.clone(), sigs.clone());
746 }
747 }
748
749 *existing_master_key = Raw::new(&existing).unwrap();
751 continue;
752 }
753 }
754
755 let known_devices = keys.device.entry(user.clone()).or_default();
759 let device_keys = known_devices
760 .get_mut(key_id)
761 .expect("trying to add a signature for a missing key");
762
763 let param: DeviceKeys = serde_json::from_str(raw_key.get()).unwrap();
764
765 let mut existing: DeviceKeys = device_keys.deserialize().unwrap();
766
767 for (uid, sigs) in existing.signatures.iter_mut() {
768 if let Some(new_sigs) = param.signatures.get(uid) {
769 sigs.extend(new_sigs.clone());
770 }
771 }
772 for (uid, sigs) in param.signatures.iter() {
773 if !existing.signatures.contains_key(uid) {
774 existing.signatures.insert(uid.clone(), sigs.clone());
775 }
776 }
777
778 *device_keys = Raw::new(&existing).unwrap();
779 }
780 }
781
782 ResponseTemplate::new(200).set_body_json(json!({
783 "failures": {}
784 }))
785 }
786}
787
788fn mock_keys_claimed_request(keys: Arc<Mutex<Keys>>) -> impl Fn(&Request) -> ResponseTemplate {
789 move |req: &Request| {
790 #[derive(Debug, serde::Deserialize)]
792 struct Parameters {
793 one_time_keys: BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, OneTimeKeyAlgorithm>>,
794 }
795
796 let params: Parameters = req.body_json().unwrap();
797
798 let mut keys = keys.lock().unwrap();
799 let known_otks = &mut keys.one_time_keys;
800
801 let mut found_one_time_keys: BTreeMap<
802 OwnedUserId,
803 BTreeMap<OwnedDeviceId, BTreeMap<OwnedOneTimeKeyId, Raw<OneTimeKey>>>,
804 > = BTreeMap::new();
805
806 for (user, requested_one_time_keys) in params.one_time_keys {
807 for device_id in requested_one_time_keys.keys() {
808 let device_id = device_id.clone();
809 let found_key = known_otks
810 .entry(user.clone())
811 .or_default()
812 .entry(device_id.clone())
813 .or_default()
814 .pop_first();
815 if let Some((id, raw_otk)) = found_key {
816 found_one_time_keys
817 .entry(user.clone())
818 .or_default()
819 .entry(device_id)
820 .or_default()
821 .insert(id, raw_otk);
822 }
823 }
824 }
825
826 ResponseTemplate::new(200).set_body_json(json!({
827 "one_time_keys" : found_one_time_keys
828 }))
829 }
830}