matrix_sdk/authentication/oauth/
cross_process.rs1use std::sync::Arc;
2
3#[cfg(feature = "e2e-encryption")]
4use matrix_sdk_base::crypto::{
5 CryptoStoreError,
6 store::{LockableCryptoStore, Store},
7};
8use matrix_sdk_common::cross_process_lock::{
9 CrossProcessLock, CrossProcessLockError, CrossProcessLockGuard,
10};
11use sha2::{Digest as _, Sha256};
12use thiserror::Error;
13use tokio::sync::{Mutex, OwnedMutexGuard};
14use tracing::trace;
15
16use crate::SessionTokens;
17
18const OIDC_SESSION_HASH_KEY: &str = "oidc_session_hash";
21
22#[derive(Clone, PartialEq, Eq)]
24struct SessionHash(Vec<u8>);
25
26impl SessionHash {
27 fn to_hex(&self) -> String {
28 const CHARS: &[char; 16] =
29 &['0', '1', '2', '3', '4', '5', '6', '7', '8', '9', 'a', 'b', 'c', 'd', 'e', 'f'];
30 let mut res = String::with_capacity(2 * self.0.len() + 2);
31 if !self.0.is_empty() {
32 res.push('0');
33 res.push('x');
34 }
35 for &c in &self.0 {
36 res.push(CHARS[(c >> 4) as usize]);
40 res.push(CHARS[(c & 0b1111) as usize]);
41 }
42 res
43 }
44}
45
46impl std::fmt::Debug for SessionHash {
47 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
48 f.debug_tuple("SessionHash").field(&self.to_hex()).finish()
49 }
50}
51
52fn compute_session_hash(tokens: &SessionTokens) -> SessionHash {
54 let mut hash = Sha256::new().chain_update(tokens.access_token.as_bytes());
55 if let Some(refresh_token) = &tokens.refresh_token {
56 hash = hash.chain_update(refresh_token.as_bytes());
57 }
58 SessionHash(hash.finalize().to_vec())
59}
60
61#[derive(Clone)]
62pub(super) struct CrossProcessRefreshManager {
63 store: Store,
64 store_lock: CrossProcessLock<LockableCryptoStore>,
65 known_session_hash: Arc<Mutex<Option<SessionHash>>>,
66}
67
68impl CrossProcessRefreshManager {
69 pub fn new(store: Store, lock: CrossProcessLock<LockableCryptoStore>) -> Self {
71 Self { store, store_lock: lock, known_session_hash: Arc::new(Mutex::new(None)) }
72 }
73
74 pub async fn spin_lock(
80 &self,
81 ) -> Result<CrossProcessRefreshLockGuard, CrossProcessRefreshLockError> {
82 trace!("Waiting for intra-process lock...");
85 let prev_hash = self.known_session_hash.clone().lock_owned().await;
86
87 trace!("Waiting for inter-process lock...");
90 let store_guard = self
91 .store_lock
92 .spin_lock(Some(60000))
93 .await
94 .map_err(|err| {
95 CrossProcessRefreshLockError::LockError(CrossProcessLockError::TryLock(Arc::new(
96 err,
97 )))
98 })?
99 .map_err(|err| CrossProcessRefreshLockError::LockError(err.into()))?;
100
101 let current_db_session_bytes = self.store.get_custom_value(OIDC_SESSION_HASH_KEY).await?;
103
104 let db_hash = current_db_session_bytes.map(SessionHash);
105
106 let hash_mismatch = match (&db_hash, &*prev_hash) {
107 (None, _) => false,
108 (Some(_), None) => true,
109 (Some(db), Some(known)) => db != known,
110 };
111
112 trace!(hash_mismatch, ?prev_hash, ?db_hash);
113
114 let guard = CrossProcessRefreshLockGuard {
115 hash_guard: prev_hash,
116 _store_guard: store_guard.into_guard(),
117 hash_mismatch,
118 db_hash,
119 store: self.store.clone(),
120 };
121
122 Ok(guard)
123 }
124
125 pub async fn restore_session(&self, tokens: &SessionTokens) {
126 let prev_tokens_hash = compute_session_hash(tokens);
127 *self.known_session_hash.lock().await = Some(prev_tokens_hash);
128 }
129
130 pub async fn on_logout(&self) -> Result<(), CrossProcessRefreshLockError> {
131 self.store
132 .remove_custom_value(OIDC_SESSION_HASH_KEY)
133 .await
134 .map_err(CrossProcessRefreshLockError::StoreError)?;
135 *self.known_session_hash.lock().await = None;
136 Ok(())
137 }
138}
139
140pub(super) struct CrossProcessRefreshLockGuard {
141 hash_guard: OwnedMutexGuard<Option<SessionHash>>,
144
145 _store_guard: CrossProcessLockGuard,
147
148 store: Store,
151
152 pub hash_mismatch: bool,
161
162 db_hash: Option<SessionHash>,
166}
167
168impl CrossProcessRefreshLockGuard {
169 fn save_in_memory(&mut self, hash: SessionHash) {
171 *self.hash_guard = Some(hash);
172 }
173
174 async fn save_in_database(
176 &self,
177 hash: &SessionHash,
178 ) -> Result<(), CrossProcessRefreshLockError> {
179 self.store.set_custom_value(OIDC_SESSION_HASH_KEY, hash.0.clone()).await?;
180 Ok(())
181 }
182
183 pub async fn save_in_memory_and_db(
187 &mut self,
188 tokens: &SessionTokens,
189 ) -> Result<(), CrossProcessRefreshLockError> {
190 let hash = compute_session_hash(tokens);
191 self.save_in_database(&hash).await?;
192 self.save_in_memory(hash);
193 Ok(())
194 }
195
196 pub async fn handle_mismatch(
199 &mut self,
200 trusted_tokens: &SessionTokens,
201 ) -> Result<(), CrossProcessRefreshLockError> {
202 let new_hash = compute_session_hash(trusted_tokens);
203 trace!("Trusted OAuth 2.0 tokens have hash {new_hash:?}; db had {:?}", self.db_hash);
204
205 if let Some(db_hash) = &self.db_hash
206 && new_hash != *db_hash
207 {
208 tracing::error!("error: DB and trusted disagree. Overriding in DB.");
212 self.save_in_database(&new_hash).await?;
213 }
214
215 self.save_in_memory(new_hash);
216 Ok(())
217 }
218}
219
220#[derive(Debug, Error)]
223pub enum CrossProcessRefreshLockError {
224 #[error(transparent)]
226 StoreError(#[from] CryptoStoreError),
227
228 #[error(transparent)]
230 LockError(#[from] CrossProcessLockError),
231
232 #[error("the previous stored hash isn't a valid integer")]
234 InvalidPreviousHash,
235
236 #[error("the cross-process lock hasn't been set up with `enable_cross_process_refresh_lock")]
238 MissingLock,
239
240 #[error(
242 "reload session callback must be set with Client::set_session_callbacks() \
243 for the cross-process lock to work"
244 )]
245 MissingReloadSession,
246
247 #[error(
249 "the cross-process lock has been set up twice with `enable_cross_process_refresh_lock`"
250 )]
251 DuplicatedLock,
252}
253
254#[cfg(all(test, feature = "e2e-encryption", feature = "sqlite", not(target_family = "wasm")))]
255mod tests {
256
257 use anyhow::Context as _;
258 use futures_util::future::join_all;
259 use matrix_sdk_base::{SessionMeta, store::RoomLoadSettings};
260 use matrix_sdk_test::async_test;
261 use ruma::{owned_device_id, owned_user_id};
262
263 use super::compute_session_hash;
264 use crate::{
265 Error,
266 authentication::oauth::cross_process::SessionHash,
267 test_utils::{
268 client::{
269 MockClientBuilder, mock_prev_session_tokens_with_refresh,
270 mock_session_tokens_with_refresh, oauth::mock_session,
271 },
272 mocks::MatrixMockServer,
273 },
274 };
275
276 #[async_test]
277 async fn test_restore_session_lock() -> Result<(), Error> {
278 let tmp_dir = tempfile::tempdir()?;
281 let client = MockClientBuilder::new(None)
282 .on_builder(|builder| builder.sqlite_store(&tmp_dir, None))
283 .unlogged()
284 .build()
285 .await;
286
287 let tokens = mock_session_tokens_with_refresh();
288
289 client.oauth().enable_cross_process_refresh_lock("test".to_owned()).await?;
290
291 client.set_session_callbacks(
292 Box::new({
293 let tokens = tokens.clone();
295 move |_| Ok(tokens.clone())
296 }),
297 Box::new(|_| panic!("save_session_callback shouldn't be called here")),
298 )?;
299
300 let session_hash = compute_session_hash(&tokens);
301 client
302 .oauth()
303 .restore_session(mock_session(tokens.clone()), RoomLoadSettings::default())
304 .await?;
305
306 assert_eq!(client.session_tokens().unwrap(), tokens);
307
308 let oauth = client.oauth();
309 let xp_manager = oauth.ctx().cross_process_token_refresh_manager.get().unwrap();
310
311 {
312 let known_session = xp_manager.known_session_hash.lock().await;
313 assert_eq!(known_session.as_ref().unwrap(), &session_hash);
314 }
315
316 {
317 let lock = xp_manager.spin_lock().await.unwrap();
318 assert!(!lock.hash_mismatch);
319 assert_eq!(lock.db_hash.unwrap(), session_hash);
320 }
321
322 Ok(())
323 }
324
325 #[async_test]
326 async fn test_finish_login() -> anyhow::Result<()> {
327 let server = MatrixMockServer::new().await;
328 server.mock_who_am_i().ok().expect(1).named("whoami").mount().await;
329
330 let tmp_dir = tempfile::tempdir()?;
331 let client = server
332 .client_builder()
333 .on_builder(|builder| builder.sqlite_store(&tmp_dir, None))
334 .registered_with_oauth()
335 .build()
336 .await;
337 let oauth = client.oauth();
338
339 oauth.enable_cross_process_refresh_lock("lock".to_owned()).await?;
341
342 let session_tokens = mock_session_tokens_with_refresh();
344 client.auth_ctx().set_session_tokens(session_tokens.clone());
345
346 oauth.load_session(owned_device_id!("D3V1C31D")).await?;
348
349 let session_meta = client.session_meta().context("should have session meta now")?;
350 assert_eq!(
351 *session_meta,
352 SessionMeta {
353 user_id: owned_user_id!("@joe:example.org"),
354 device_id: owned_device_id!("D3V1C31D")
355 }
356 );
357
358 {
359 let xp_manager =
362 oauth.ctx().cross_process_token_refresh_manager.get().context("must have lock")?;
363 let guard = xp_manager.spin_lock().await?;
364 let actual_hash = compute_session_hash(&session_tokens);
365 assert_eq!(guard.db_hash.as_ref(), Some(&actual_hash));
366 assert_eq!(guard.hash_guard.as_ref(), Some(&actual_hash));
367 assert!(!guard.hash_mismatch);
368 }
369
370 Ok(())
371 }
372
373 #[async_test]
374 async fn test_refresh_access_token_twice() -> anyhow::Result<()> {
375 let server = MatrixMockServer::new().await;
380
381 let oauth_server = server.oauth();
382 oauth_server.mock_server_metadata().ok().expect(1..).named("server_metadata").mount().await;
383 oauth_server.mock_token().ok().expect(1).named("token").mount().await;
384
385 let tmp_dir = tempfile::tempdir()?;
386 let client = server
387 .client_builder()
388 .on_builder(|builder| builder.sqlite_store(&tmp_dir, None))
389 .unlogged()
390 .build()
391 .await;
392 let oauth = client.oauth();
393
394 let next_tokens = mock_session_tokens_with_refresh();
395
396 oauth.enable_cross_process_refresh_lock("lock".to_owned()).await?;
398
399 oauth
401 .restore_session(
402 mock_session(mock_prev_session_tokens_with_refresh()),
403 RoomLoadSettings::default(),
404 )
405 .await?;
406
407 for result in join_all([oauth.refresh_access_token(), oauth.refresh_access_token()]).await {
409 result?;
410 }
411
412 {
413 let xp_manager =
416 oauth.ctx().cross_process_token_refresh_manager.get().context("must have lock")?;
417 let guard = xp_manager.spin_lock().await?;
418 let actual_hash = compute_session_hash(&next_tokens);
419 assert_eq!(guard.db_hash.as_ref(), Some(&actual_hash));
420 assert_eq!(guard.hash_guard.as_ref(), Some(&actual_hash));
421 assert!(!guard.hash_mismatch);
422 }
423
424 Ok(())
425 }
426
427 #[async_test]
428 async fn test_cross_process_concurrent_refresh() -> anyhow::Result<()> {
429 let server = MatrixMockServer::new().await;
430
431 let oauth_server = server.oauth();
432 oauth_server.mock_server_metadata().ok().expect(1..).named("server_metadata").mount().await;
433 oauth_server.mock_token().ok().expect(1).named("token").mount().await;
434
435 let prev_tokens = mock_prev_session_tokens_with_refresh();
436 let next_tokens = mock_session_tokens_with_refresh();
437
438 let tmp_dir = tempfile::tempdir()?;
440 let client = server
441 .client_builder()
442 .on_builder(|builder| builder.sqlite_store(&tmp_dir, None))
443 .unlogged()
444 .build()
445 .await;
446
447 let oauth = client.oauth();
448 oauth.enable_cross_process_refresh_lock("client1".to_owned()).await?;
449
450 oauth
451 .restore_session(mock_session(prev_tokens.clone()), RoomLoadSettings::default())
452 .await?;
453
454 let unrestored_client = server
457 .client_builder()
458 .on_builder(|builder| builder.sqlite_store(&tmp_dir, None))
459 .unlogged()
460 .build()
461 .await;
462 let unrestored_oauth = unrestored_client.oauth();
463 unrestored_oauth.enable_cross_process_refresh_lock("unrestored_client".to_owned()).await?;
464
465 {
466 let client3 = server
469 .client_builder()
470 .on_builder(|builder| builder.sqlite_store(&tmp_dir, None))
471 .unlogged()
472 .build()
473 .await;
474
475 let oauth3 = client3.oauth();
476 oauth3.enable_cross_process_refresh_lock("client3".to_owned()).await?;
477 oauth3
478 .restore_session(mock_session(prev_tokens.clone()), RoomLoadSettings::default())
479 .await?;
480
481 oauth3.refresh_access_token().await?;
484
485 assert_eq!(client3.session_tokens(), Some(next_tokens.clone()));
486
487 let xp_manager =
490 oauth3.ctx().cross_process_token_refresh_manager.get().context("must have lock")?;
491 let guard = xp_manager.spin_lock().await?;
492 let actual_hash = compute_session_hash(&next_tokens);
493 assert_eq!(guard.db_hash.as_ref(), Some(&actual_hash));
494 assert_eq!(guard.hash_guard.as_ref(), Some(&actual_hash));
495 assert!(!guard.hash_mismatch);
496 }
497
498 {
499 let oauth = unrestored_oauth;
502
503 unrestored_client.set_session_callbacks(
504 Box::new({
505 let tokens = next_tokens.clone();
507 move |_| Ok(tokens.clone())
508 }),
509 Box::new(|_| panic!("save_session_callback shouldn't be called here")),
510 )?;
511
512 oauth
513 .restore_session(mock_session(prev_tokens.clone()), RoomLoadSettings::default())
514 .await?;
515
516 let xp_manager =
518 oauth.ctx().cross_process_token_refresh_manager.get().context("must have lock")?;
519 let guard = xp_manager.spin_lock().await?;
520 let next_hash = compute_session_hash(&next_tokens);
521 assert_eq!(guard.db_hash.as_ref(), Some(&next_hash));
522 assert_eq!(guard.hash_guard.as_ref(), Some(&next_hash));
523 assert!(!guard.hash_mismatch);
524
525 drop(oauth);
526 drop(unrestored_client);
527 }
528
529 {
530 let xp_manager =
533 oauth.ctx().cross_process_token_refresh_manager.get().context("must have lock")?;
534 let guard = xp_manager.spin_lock().await?;
535 let previous_hash = compute_session_hash(&prev_tokens);
536 let next_hash = compute_session_hash(&next_tokens);
537 assert_eq!(guard.db_hash, Some(next_hash));
538 assert_eq!(guard.hash_guard.as_ref(), Some(&previous_hash));
539 assert!(guard.hash_mismatch);
540 }
541
542 client.set_session_callbacks(
543 Box::new({
544 let tokens = next_tokens.clone();
546 move |_| Ok(tokens.clone())
547 }),
548 Box::new(|_| panic!("save_session_callback shouldn't be called here")),
549 )?;
550
551 oauth.refresh_access_token().await?;
552
553 {
554 let xp_manager =
556 oauth.ctx().cross_process_token_refresh_manager.get().context("must have lock")?;
557 let guard = xp_manager.spin_lock().await?;
558 let actual_hash = compute_session_hash(&next_tokens);
559 assert_eq!(guard.db_hash.as_ref(), Some(&actual_hash));
560 assert_eq!(guard.hash_guard.as_ref(), Some(&actual_hash));
561 assert!(!guard.hash_mismatch);
562 }
563
564 Ok(())
565 }
566
567 #[async_test]
568 async fn test_logout() -> anyhow::Result<()> {
569 let server = MatrixMockServer::new().await;
570
571 let oauth_server = server.oauth();
572 oauth_server
573 .mock_server_metadata()
574 .ok_https()
575 .expect(1..)
576 .named("server_metadata")
577 .mount()
578 .await;
579 oauth_server.mock_revocation().ok().expect(1).named("revocation").mount().await;
580
581 let tmp_dir = tempfile::tempdir()?;
582 let client = server
583 .client_builder()
584 .on_builder(|builder| builder.sqlite_store(&tmp_dir, None))
585 .unlogged()
586 .build()
587 .await;
588 let oauth = client.oauth().insecure_rewrite_https_to_http();
589
590 oauth.enable_cross_process_refresh_lock("lock".to_owned()).await?;
592
593 let tokens = mock_session_tokens_with_refresh();
595 oauth.restore_session(mock_session(tokens.clone()), RoomLoadSettings::default()).await?;
596
597 oauth.logout().await.unwrap();
598
599 {
600 let xp_manager =
603 oauth.ctx().cross_process_token_refresh_manager.get().context("must have lock")?;
604 let guard = xp_manager.spin_lock().await?;
605 assert!(guard.db_hash.is_none());
606 assert!(guard.hash_guard.is_none());
607 assert!(!guard.hash_mismatch);
608 }
609
610 Ok(())
611 }
612
613 #[test]
614 fn test_session_hash_to_hex() {
615 let hash = SessionHash(vec![]);
616 assert_eq!(hash.to_hex(), "");
617
618 let hash = SessionHash(vec![0x13, 0x37, 0x42, 0xde, 0xad, 0xca, 0xfe]);
619 assert_eq!(hash.to_hex(), "0x133742deadcafe");
620 }
621
622 #[async_test]
634 async fn test_refresh_interrupted_by_suspension_does_not_sign_out() {
635 use std::{thread, time::Duration};
636
637 let server = MatrixMockServer::new().await;
638 let oauth_server = server.oauth();
639
640 oauth_server
646 .mock_token()
647 .ok_with_tokens("1234", "ZYXWV") .mock_once()
649 .with_priority(1)
650 .mount()
651 .await;
652 oauth_server.mock_token().invalid_grant().with_priority(2).mount().await;
653
654 oauth_server
658 .mock_server_metadata()
659 .with_delay(Duration::from_secs(1))
660 .ok()
661 .mock_once()
662 .with_priority(1)
663 .mount()
664 .await;
665 oauth_server.mock_server_metadata().ok().with_priority(2).mount().await;
666
667 let tmp_dir = tempfile::tempdir().unwrap();
670
671 let app = server
672 .client_builder()
673 .on_builder(|b| b.sqlite_store(&tmp_dir, None))
674 .unlogged()
675 .build()
676 .await;
677 app.oauth().enable_cross_process_refresh_lock("app".to_owned()).await.unwrap();
678 app.oauth()
679 .restore_session(
680 mock_session(mock_prev_session_tokens_with_refresh()),
681 RoomLoadSettings::default(),
682 )
683 .await
684 .unwrap();
685 app.set_session_callbacks(
686 Box::new(|_| Ok(mock_session_tokens_with_refresh())),
687 Box::new(|_| Ok(())),
688 )
689 .unwrap();
690
691 let nse = server
692 .client_builder()
693 .on_builder(|b| b.sqlite_store(&tmp_dir, None))
694 .unlogged()
695 .build()
696 .await;
697 nse.oauth().enable_cross_process_refresh_lock("nse".to_owned()).await.unwrap();
698 nse.oauth()
699 .restore_session(
700 mock_session(mock_prev_session_tokens_with_refresh()),
701 RoomLoadSettings::default(),
702 )
703 .await
704 .unwrap();
705 nse.set_session_callbacks(
706 Box::new(|_| Ok(mock_session_tokens_with_refresh())),
707 Box::new(|_| Ok(())),
708 )
709 .unwrap();
710
711 let app_oauth = app.oauth();
714 let app_refresh = tokio::spawn(async move { app_oauth.refresh_access_token().await });
715
716 let mut waited = Duration::ZERO;
720 while !server
721 .received_requests()
722 .await
723 .unwrap_or_default()
724 .iter()
725 .any(|request| request.url.path().contains("auth_metadata"))
726 {
727 assert!(waited < Duration::from_secs(5), "the app never asked for the metadata");
728 tokio::time::sleep(Duration::from_millis(10)).await;
729 waited += Duration::from_millis(10);
730 }
731
732 thread::sleep(Duration::from_millis(700));
736
737 nse.oauth().refresh_access_token().await.unwrap();
740 assert_eq!(nse.session_tokens(), Some(mock_session_tokens_with_refresh()));
741
742 let app_refresh = app_refresh.await.expect("the app refresh task shouldn't panic");
746
747 assert!(
748 app_refresh.is_ok(),
749 "the app was signed out after the NSE rotated the refresh token: {app_refresh:?}"
750 );
751 }
752}