Skip to main content

matrix_sdk/test_utils/mocks/
encryption.rs

1// Copyright 2024 The Matrix.org Foundation C.I.C.
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Helpers to mock a server that supports the main crypto API and have a client
16//! automatically connected to that server, for the purpose of integration
17//! tests.
18use 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
54/// Stores pending to-device messages for each user and device. To be used with
55/// [`MatrixMockServer::capture_put_to_device_traffic`].
56pub type PendingToDeviceMessages =
57    BTreeMap<OwnedUserId, BTreeMap<OwnedDeviceId, Vec<Raw<AnyToDeviceEvent>>>>;
58
59/// Extends the `MatrixMockServer` with useful methods to help mocking matrix
60/// crypto API and perform integration test with encryption.
61///
62/// It implements mock endpoints for the `keys/upload`, will store the uploaded
63/// devices and serves them back for incoming `keys/query`. It is also storing
64/// and claiming one-time-keys, allowing to set up working olm sessions.
65///
66/// Adds some helpers like `exhaust_one_time_keys` that allows to simulate a
67/// client running out of otks. More can be added if needed later.
68///
69/// It works like this:
70///
71/// - Start by creating the mock server like this [`MatrixMockServer::new`].
72/// - Then mock the crypto API endpoints
73///   [`MatrixMockServer::mock_crypto_endpoints_preset`].
74/// - Create your test client using
75///   [`MatrixMockServer::client_builder_for_crypto_end_to_end`], this is
76///   important as it will set up an access token that will allow to know what
77///   client is doing what request.
78///
79/// The [`MatrixMockServer::set_up_alice_and_bob_for_encryption`] will set up
80/// two olm machines aware of each other and ready to communicate.
81impl MatrixMockServer {
82    /// Creates a new [`MockClientBuilder`] configured to use this server and
83    /// suitable for usage of the crypto API end points. Will create a specific
84    /// access token and some mapping to the associated user_id.
85    pub fn client_builder_for_crypto_end_to_end(
86        &self,
87        user_id: &UserId,
88        device_id: &DeviceId,
89    ) -> MockClientBuilder {
90        // Create an access token and store the token to user_id mapping
91        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    /// Makes the server forget about all the one-time-keys for that device.
108    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    /// Ensure that the given clients are aware of each others public
115    /// identities.
116    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        // Have Alice track Bob, so she queries his keys later.
124        alice.update_tracked_users_for_testing([bob_user_id]).instrument(alice_span.clone()).await;
125
126        // let bob be aware of Alice keys in order to be able to decrypt custom
127        // to-device (the device keys check are deferred for `m.room.key` so
128        // this is not needed for sending room messages for example).
129        bob.update_tracked_users_for_testing([alice_user_id]).instrument(bob_span.clone()).await;
130
131        // Have Alice and Bob upload their signed device keys.
132        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        // Run a sync so we do send outgoing requests, including the /keys/query
136        // for getting bob's identity.
137        self.mock_sync().ok_and_run(alice, |_x| {}).instrument(alice_span).await;
138    }
139
140    /// Utility to properly setup two clients. These two clients will know about
141    /// each others (alice will have downloaded bob device keys).
142    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    /// Creates a third client for e2e tests.
162    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        // Let carl upload it's device keys.
170        self.mock_sync().ok_and_run(&carl, |_| {}).await;
171
172        // Have Alice track Carl, so she queries his keys later.
173        alice.update_tracked_users_for_testing([carl.user_id().unwrap()]).await;
174
175        // Have Bob track Carl, so she queries his keys later.
176        bob.update_tracked_users_for_testing([carl.user_id().unwrap()]).await;
177
178        // Have Alice and Bob upload their signed device keys, and download
179        // Carl's keys.
180        {
181            self.mock_sync().ok_and_run(alice, |_| {}).await;
182            self.mock_sync().ok_and_run(bob, |_| {}).await;
183        }
184
185        // Let carl be aware of Alice and Bob keys.
186        carl.update_tracked_users_for_testing([alice.user_id().unwrap(), bob.user_id().unwrap()])
187            .await;
188
189        // A last sync for carl to get the keys.
190        self.mock_sync().ok_and_run(alice, |_| {}).await;
191
192        carl
193    }
194
195    /// Creates a new device and returns a new client for it. The new and old
196    /// clients will be aware of each other.
197    ///
198    /// # Arguments
199    ///
200    /// - `existing_client` - The original client for which a new device will be
201    ///   created
202    /// - `device_id` - The device ID to use for the new client
203    /// - `clients_to_update` - A vector of client references that should be
204    ///   notified about the new device. These clients will receive a device
205    ///   list change notification during their next sync.
206    ///
207    /// # Returns
208    ///
209    /// Returns the newly created client instance configured for the new device.
210    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        // sync the keys
223        self.mock_sync().ok_and_run(&new_client, |_| {}).await;
224
225        // Notify existing device of a change
226        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    /// Mock up the various crypto API so that it can serve back keys when
244    /// needed
245    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    /// Creates a response handler for mocking encrypted to-device message
281    /// requests.
282    ///
283    /// This function creates a response handler that captures encrypted
284    /// to-device messages sent via the `/sendToDevice` endpoint.
285    ///
286    /// # Arguments
287    ///
288    /// - `sender` - The user ID of the message sender
289    ///
290    /// # Returns
291    ///
292    /// Returns a tuple containing:
293    ///
294    /// - A `MockGuard` the end-point mock is scoped to this guard
295    /// - A `Future` that resolves to a `Raw<EncryptedToDeviceEvent>>`
296    ///   containing the captured encrypted to-device message.
297    ///
298    /// # Examples
299    ///
300    /// ```rust
301    /// # use ruma::{ device_id,  user_id, serde::Raw};
302    /// # use serde_json::json;
303    ///
304    /// # use matrix_sdk_test::async_test;
305    /// # use matrix_sdk::test_utils::mocks::MatrixMockServer;
306    /// #
307    /// #[async_test]
308    /// async fn test_mock_capture_put_to_device() {
309    ///     let server = MatrixMockServer::new().await;
310    ///     server.mock_crypto_endpoints_preset().await;
311    ///
312    ///     let (alice, bob) = server.set_up_alice_and_bob_for_encryption().await;
313    ///     let bob_user_id = bob.user_id().unwrap();
314    ///     let bob_device_id = bob.device_id().unwrap();
315    ///
316    ///     // From the point of view of Alice, Bob now has a device.
317    ///     let alice_bob_device = alice
318    ///         .encryption()
319    ///         .get_device(bob_user_id, bob_device_id)
320    ///         .await
321    ///         .unwrap()
322    ///         .expect("alice sees bob's device");
323    ///
324    ///     let content_raw = Raw::new(&json!({ /*...*/ })).unwrap().cast();
325    ///
326    ///     // Set up the mock to capture encrypted to-device messages
327    ///     let (guard, captured) =
328    ///         server.mock_capture_put_to_device(alice.user_id().unwrap()).await;
329    ///
330    ///     alice
331    ///         .encryption()
332    ///         .encrypt_and_send_raw_to_device(
333    ///             vec![&alice_bob_device],
334    ///             "call.keys",
335    ///             content_raw,
336    ///         )
337    ///         .await
338    ///         .unwrap();
339    ///
340    ///     // this is the captured event as sent by alice!
341    ///     let sent_event = captured.await;
342    ///     drop(guard);
343    /// }
344    /// ```
345    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            // Should be called once
383            .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    /// Captures a to-device message when it is sent to the mock server and then
395    /// injects it into the recipient's sync response.
396    ///
397    /// This is a utility function that combines capturing an encrypted
398    /// to-device message and delivering it to the recipient through a sync
399    /// response. It's useful for testing end-to-end encryption scenarios where
400    /// you need to verify message delivery and processing.
401    ///
402    /// # Arguments
403    ///
404    /// - `sender_user_id` - The user ID of the message sender
405    /// - `recipient` - The client that will receive the message through sync
406    ///
407    /// # Returns
408    ///
409    /// Returns a `Future` that will resolve when the captured event has been
410    /// fed back down the recipient sync.
411    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    /// Utility to capture all the `/toDevice` upload traffic and store it in a
432    /// queue to be later used with
433    /// [`MatrixMockServer::sync_back_pending_to_device_messages`].
434    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                // Access the captured groups from the path
453                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    /// Sync the pending to-device messages for this client.
487    ///
488    /// To be used in connection with
489    /// [`MatrixMockServer::capture_put_to_device_traffic`] that is capturing
490    /// the traffic.
491    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
515/// Intercepts a `/keys/query` request and mock its results as returned by an
516/// actual homeserver.
517///
518/// Supports filtering by user id, or no filters at all.
519fn 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
558/// Intercepts a `/keys/upload` query and mocks the behavior it would have on a
559/// real homeserver.
560///
561/// Inserts all the `DeviceKeys` into `Keys::device_keys`, or if already present
562/// in this mapping, only merge the signatures.
563fn 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        // Get the user
583        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            // if known_devices.contains(&key_id) {
592            let mut keys = keys.lock().unwrap();
593            let devices = keys.device.entry(new_device_keys.user_id.clone()).or_default();
594
595            // Either merge signatures if an entry is already present, or insert
596            // a new one.
597            if let Some(device_keys) = devices.get_mut(&key_id) {
598                let mut existing = device_keys.deserialize().unwrap();
599
600                // Merge signatures.
601                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            // We need a trick to find out what userId|device this OTK is for.
622            // This is not part of the payload, a real server uses the access
623            // token(?) Let's look at the signatures to find out
624            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                        // Ignore this old algorithm,
648                    }
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
663/// Mocks a `/keys/device_signing/upload` request for bootstrapping
664/// cross-signing.
665///
666/// Assumes (and asserts) all keys are updated at the same time.
667///
668/// Saves all the different cross-signing keys into their respective fields of
669/// `Keys`.
670fn mock_keys_device_signing_upload(
671    keys: Arc<Mutex<Keys>>,
672) -> impl Fn(&Request) -> ResponseTemplate {
673    move |req: &Request| {
674        // Accept all cross-signing setups by default.
675        #[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
711/// Mocks a `/keys/signatures/upload` request.
712///
713/// Supports merging signatures for master keys or devices keys.
714fn 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                // Try to find a field in keys.master.
727                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                        // Update in map.
750                        *existing_master_key = Raw::new(&existing).unwrap();
751                        continue;
752                    }
753                }
754
755                // Otherwise, try to find a field in keys.device. Either merge
756                // signatures if an entry is already present, or insert a new
757                // entry.
758                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        // Accept all cross-signing setups by default.
791        #[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}