From cc2d8b1bfbd3b15392ab3ebcf1e9b4b74b7f76d8 Mon Sep 17 00:00:00 2001 From: Erwan Leboucher Date: Tue, 28 Jul 2026 12:40:51 +0200 Subject: [PATCH] fix: recover media auth after token refresh --- src-tauri/src/network/media_protocol.rs | 553 ++++++++++++++++++++++-- src/client/oidcTokenRefresher.test.ts | 31 ++ src/client/oidcTokenRefresher.ts | 63 ++- src/serviceWorkerBootstrap.test.ts | 55 ++- src/serviceWorkerBootstrap.ts | 29 +- src/sw-media-auth-recovery.test.ts | 28 ++ 6 files changed, 698 insertions(+), 61 deletions(-) diff --git a/src-tauri/src/network/media_protocol.rs b/src-tauri/src/network/media_protocol.rs index 536e2344ce..c3fe293cbe 100644 --- a/src-tauri/src/network/media_protocol.rs +++ b/src-tauri/src/network/media_protocol.rs @@ -4,7 +4,7 @@ use std::{ io::Read, path::{Path, PathBuf}, sync::{ - atomic::{AtomicBool, Ordering}, + atomic::{AtomicBool, AtomicU64, Ordering}, Arc, Mutex, OnceLock, RwLock, Weak, }, time::Duration, @@ -29,6 +29,7 @@ use tokio::{ }; pub const MEDIA_URI_SCHEME: &str = "sable-media"; +const MEDIA_SESSION_MARKER: &str = "__sable_media_session"; const MEDIA_PATH_PREFIXES: [&str; 2] = ["/_matrix/media/", "/_matrix/client/v1/media/"]; // How the webview spells this protocol: `sable-media://` on iOS/macOS, and @@ -47,6 +48,7 @@ const MAX_CONCURRENT_DOWNLOAD_REQUESTS: usize = 6; // The frontend mounts (and starts requesting media) before it hands us the session, so a request // may arrive first. `` never retries, so waiting beats answering 503. const SESSION_WAIT: Duration = Duration::from_secs(5); +const SESSION_REFRESH_WAIT: Duration = Duration::from_millis(2500); const MAX_CACHE_BYTES: u64 = 512 * 1024 * 1024; // Short: a 4xx stops being true once an unreachable remote server comes back. const NEGATIVE_CACHE_TTL: Duration = Duration::from_secs(600); @@ -61,6 +63,7 @@ pub struct MediaSessionState { inner: RwLock>, session_ready: Notify, session_ever_set: AtomicBool, + session_generation: AtomicU64, encryption: RwLock>, client: OnceLock, thumbnail_semaphore: Semaphore, @@ -75,6 +78,7 @@ impl Default for MediaSessionState { inner: RwLock::new(None), session_ready: Notify::new(), session_ever_set: AtomicBool::new(false), + session_generation: AtomicU64::new(0), encryption: RwLock::new(HashMap::new()), client: OnceLock::new(), thumbnail_semaphore: Semaphore::new(MAX_CONCURRENT_THUMBNAIL_REQUESTS), @@ -114,6 +118,68 @@ impl MediaSessionState { } } + async fn wait_for_newer_session(&self, previous: &MediaSession) -> Option { + self.wait_for_newer_session_with_timeout(previous, SESSION_REFRESH_WAIT) + .await + } + + async fn wait_for_newer_session_with_timeout( + &self, + previous: &MediaSession, + wait: Duration, + ) -> Option { + let deadline = Instant::now() + wait; + loop { + let mut notified = std::pin::pin!(self.session_ready.notified()); + // Register before re-checking, otherwise a session arriving in between is missed. + notified.as_mut().enable(); + match self.session() { + Some(session) if session.generation > previous.generation => { + if session.origin != previous.origin || session.scope != previous.scope { + return None; + } + if session.token != previous.token { + return Some(session); + } + } + None if self.session_generation.load(Ordering::Acquire) > previous.generation => { + return None; + } + _ => {} + } + if timeout_at(deadline, notified).await.is_err() { + return None; + } + } + } + + fn set_session(&self, mut session: MediaSession) -> Result<(), String> { + let mut guard = self + .inner + .write() + .map_err(|_| "media session lock poisoned".to_string())?; + session.generation = self.session_generation.fetch_add(1, Ordering::AcqRel) + 1; + *guard = Some(session); + drop(guard); + self.forget_client_errors(); + self.session_ever_set.store(true, Ordering::Release); + self.session_ready.notify_waiters(); + Ok(()) + } + + fn clear_session(&self) -> Result<(), String> { + let mut guard = self + .inner + .write() + .map_err(|_| "media session lock poisoned".to_string())?; + *guard = None; + self.session_generation.fetch_add(1, Ordering::AcqRel); + drop(guard); + self.forget_client_errors(); + self.session_ready.notify_waiters(); + Ok(()) + } + // Shared across requests so the connection pool and TLS sessions stay warm. fn client(&self) -> Client { self.client @@ -194,6 +260,7 @@ struct MediaSession { // Cache key input. The Matrix user ID, not `token`, which rotates on every OIDC // refresh and would orphan the whole on-disk cache. scope: String, + generation: u64, } #[derive(Clone)] @@ -221,22 +288,12 @@ pub fn set_media_session( .filter(|value| !value.is_empty()) .unwrap_or_else(|| origin.clone()); - { - let mut guard = state - .inner - .write() - .map_err(|_| "media session lock poisoned".to_string())?; - *guard = Some(MediaSession { - origin, - token, - scope, - }); - } - - state.forget_client_errors(); - state.session_ever_set.store(true, Ordering::Release); - state.session_ready.notify_waiters(); - Ok(()) + state.set_session(MediaSession { + origin, + token, + scope, + generation: 0, + }) } #[tauri::command] @@ -244,19 +301,16 @@ pub fn clear_media_session( app: AppHandle, state: tauri::State<'_, MediaSessionState>, ) { - if let Ok(mut guard) = state.inner.write() { - *guard = None; - } - if let Ok(mut guard) = state.encryption.write() { - guard.clear(); - } - state.forget_client_errors(); + let _ = state.clear_session(); if let Ok(dir) = cache_dir(&app) { let _ = fs::remove_dir_all(dir); } if let Ok(temp_dir) = temp_cache_dir(&app) { let _ = fs::remove_dir_all(temp_dir); } + if let Ok(mut guard) = state.encryption.write() { + guard.clear(); + } } #[tauri::command] @@ -341,6 +395,44 @@ fn normalize_encryption_key(url: &str) -> String { .unwrap_or_else(|_| url.to_string()) } +fn media_session_marker(uri: &Uri) -> Result, StatusCode> { + let Some(query) = uri.query() else { + return Ok(None); + }; + + let mut marker = None; + for component in query.split('&') { + let (raw_key, raw_value) = component.split_once('=').unwrap_or((component, "")); + let key = percent_encoding::percent_decode_str(raw_key) + .decode_utf8() + .map_err(|_| StatusCode::BAD_REQUEST)?; + if key == MEDIA_SESSION_MARKER { + if marker.is_some() { + return Err(StatusCode::BAD_REQUEST); + } + marker = Some( + percent_encoding::percent_decode_str(raw_value) + .decode_utf8() + .map_err(|_| StatusCode::BAD_REQUEST)? + .into_owned(), + ); + } + } + + Ok(marker) +} + +fn session_marker_matches(uri: &Uri, scope: &str) -> Result { + Ok(media_session_marker(uri)?.is_none_or(|marker| marker == scope)) +} + +fn should_retry_with_session(failed: &MediaSession, updated: &MediaSession) -> bool { + updated.generation != failed.generation + && updated.token != failed.token + && updated.origin == failed.origin + && updated.scope == failed.scope +} + /// The app's own webview origins, which vary by platform and scheme. fn is_webview_origin(origin: &str) -> bool { matches!( @@ -415,6 +507,9 @@ async fn handle_request( { return Err(StatusCode::FORBIDDEN); } + if !session_marker_matches(&uri, &session.scope)? { + return Err(StatusCode::FORBIDDEN); + } let key = cache_key(&session.scope, &target); if let Some(status) = state.recent_client_error(&key) { @@ -612,6 +707,23 @@ async fn fetch_and_cache( .await .map_err(|_| StatusCode::BAD_GATEWAY)?; + if matches!( + upstream.status(), + StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN + ) { + if let Some(updated) = state.wait_for_newer_session(session).await { + if should_retry_with_session(session, &updated) { + upstream = state + .client() + .get(media_url.clone()) + .header(AUTHORIZATION, format!("Bearer {}", updated.token)) + .send() + .await + .map_err(|_| StatusCode::BAD_GATEWAY)?; + } + } + } + if !upstream.status().is_success() { return Err( StatusCode::from_u16(upstream.status().as_u16()).unwrap_or(StatusCode::BAD_GATEWAY) @@ -1069,11 +1181,123 @@ fn cache_key(session_token: &str, url: &str) -> String { #[cfg(test)] mod tests { - use std::sync::Arc; + use std::{ + fs, + io::{Read, Write}, + net::TcpListener, + sync::{atomic::AtomicU64, mpsc, Arc}, + thread, + time::Duration, + }; + + use tauri::http::{header, StatusCode, Uri}; - use tauri::http::{header, StatusCode}; + use super::{ + cache_key, ok_response, session_marker_matches, should_retry_with_session, MediaSession, + MediaSessionState, Url, + }; - use super::{cache_key, ok_response, MediaSessionState}; + static TEST_CACHE_ID: AtomicU64 = AtomicU64::new(0); + + fn test_session(origin: &str, token: &str, scope: &str, generation: u64) -> MediaSession { + MediaSession { + origin: origin.to_owned(), + token: token.to_owned(), + scope: scope.to_owned(), + generation, + } + } + + fn start_upstream( + statuses: Vec<(u16, &'static str)>, + ) -> (Url, mpsc::Receiver, thread::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = listener.local_addr().unwrap(); + let (auth_sender, auth_receiver) = mpsc::channel(); + let server = thread::spawn(move || { + for (status, body) in statuses { + let (mut stream, _) = listener.accept().unwrap(); + let mut request = Vec::new(); + let mut chunk = [0_u8; 1024]; + while !request.windows(4).any(|window| window == b"\r\n\r\n") { + let Ok(read) = stream.read(&mut chunk) else { + break; + }; + if read == 0 { + break; + } + request.extend_from_slice(&chunk[..read]); + } + let authorization = String::from_utf8_lossy(&request) + .lines() + .find_map(|line| { + line.strip_prefix("Authorization: ") + .or_else(|| line.strip_prefix("authorization: ")) + }) + .unwrap_or_default() + .to_owned(); + auth_sender.send(authorization).unwrap(); + let reason = match status { + 200 => "OK", + 403 => "Forbidden", + _ => "Unauthorized", + }; + let response = format!( + "HTTP/1.1 {status} {reason}\r\nContent-Type: image/png\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}", + body.len() + ); + stream.write_all(response.as_bytes()).unwrap(); + } + }); + ( + Url::parse(&format!( + "http://{address}/_matrix/client/v1/media/download/matrix.org/id" + )) + .unwrap(), + auth_receiver, + server, + ) + } + + async fn fetch_test( + state: &MediaSessionState, + session: MediaSession, + media_url: Url, + ) -> Result<(String, Option>>, std::path::PathBuf), StatusCode> { + let root = std::env::temp_dir().join(format!( + "sable-media-test-{}-{}", + std::process::id(), + TEST_CACHE_ID.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + )); + let temp = root.join("temp"); + let key = cache_key(&session.scope, media_url.as_str()); + let result = super::ensure_cached_with_limits( + state, + &session, + &key, + media_url, + root.clone(), + temp, + 1024 * 1024, + 1024 * 1024, + ) + .await; + fs::remove_dir_all(root).ok(); + result + } + + async fn receive_auth(receiver: &mpsc::Receiver) -> String { + tokio::time::timeout(Duration::from_secs(1), async { + loop { + if let Ok(auth) = receiver.try_recv() { + return auth; + } + tokio::task::yield_now().await; + } + }) + .await + .unwrap() + } #[test] fn cache_miss_gates_are_reused_and_expired_entries_are_cleaned() { @@ -1101,18 +1325,14 @@ mod tests { let writer = Arc::clone(&state); tokio::spawn(async move { tokio::time::sleep(std::time::Duration::from_millis(50)).await; - { - let mut guard = writer.inner.write().unwrap(); - *guard = Some(super::MediaSession { + writer + .set_session(super::MediaSession { origin: "https://matrix.example.org".into(), token: "token".into(), scope: "@a:example.org".into(), - }); - } - writer - .session_ever_set - .store(true, std::sync::atomic::Ordering::Release); - writer.session_ready.notify_waiters(); + generation: 0, + }) + .unwrap(); }); let session = @@ -1122,6 +1342,267 @@ mod tests { assert_eq!(session.map(|s| s.scope), Some("@a:example.org".to_string())); } + #[tokio::test] + async fn retries_once_after_a_same_scope_token_update() { + let (media_url, auth_receiver, server) = start_upstream(vec![(401, ""), (200, "data")]); + let state = Arc::new(MediaSessionState::default()); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "old-token", + "@a:example.org", + 0, + )) + .unwrap(); + let failed_session = state.session().unwrap(); + let fetch_state = Arc::clone(&state); + let fetch_url = media_url.clone(); + let fetch = + tokio::spawn(async move { fetch_test(&fetch_state, failed_session, fetch_url).await }); + + assert_eq!(receive_auth(&auth_receiver).await, "Bearer old-token"); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "new-token", + "@a:example.org", + 0, + )) + .unwrap(); + + assert!(fetch.await.unwrap().is_ok()); + assert_eq!(receive_auth(&auth_receiver).await, "Bearer new-token"); + server.join().unwrap(); + } + + #[tokio::test] + async fn retries_once_after_same_token_then_changed_token_updates() { + let (media_url, auth_receiver, server) = start_upstream(vec![(401, ""), (200, "data")]); + let state = Arc::new(MediaSessionState::default()); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "same-token", + "@a:example.org", + 0, + )) + .unwrap(); + let failed_session = state.session().unwrap(); + let fetch_state = Arc::clone(&state); + let fetch_url = media_url.clone(); + let fetch = + tokio::spawn(async move { fetch_test(&fetch_state, failed_session, fetch_url).await }); + + assert_eq!(receive_auth(&auth_receiver).await, "Bearer same-token"); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "same-token", + "@a:example.org", + 0, + )) + .unwrap(); + tokio::time::sleep(Duration::from_millis(20)).await; + assert!(auth_receiver.try_recv().is_err()); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "changed-token", + "@a:example.org", + 0, + )) + .unwrap(); + + assert!(fetch.await.unwrap().is_ok()); + assert_eq!(receive_auth(&auth_receiver).await, "Bearer changed-token"); + server.join().unwrap(); + } + + #[tokio::test] + async fn times_out_without_a_session_update() { + let state = MediaSessionState::default(); + let session = test_session("https://matrix.example.org", "token", "@a:example.org", 1); + assert!(state + .wait_for_newer_session_with_timeout(&session, Duration::from_millis(20)) + .await + .is_none()); + } + + #[tokio::test] + async fn a_failed_retry_does_not_loop() { + let (media_url, auth_receiver, server) = start_upstream(vec![(401, ""), (403, "")]); + let state = Arc::new(MediaSessionState::default()); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "old-token", + "@a:example.org", + 0, + )) + .unwrap(); + let failed_session = state.session().unwrap(); + let fetch_state = Arc::clone(&state); + let fetch_url = media_url.clone(); + let fetch = + tokio::spawn(async move { fetch_test(&fetch_state, failed_session, fetch_url).await }); + + assert_eq!(receive_auth(&auth_receiver).await, "Bearer old-token"); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "new-token", + "@a:example.org", + 0, + )) + .unwrap(); + + assert_eq!(fetch.await.unwrap().unwrap_err(), StatusCode::FORBIDDEN); + assert_eq!(receive_auth(&auth_receiver).await, "Bearer new-token"); + server.join().unwrap(); + } + + #[tokio::test] + async fn clear_while_waiting_terminates_without_a_retry() { + let (media_url, auth_receiver, server) = start_upstream(vec![(401, "")]); + let state = Arc::new(MediaSessionState::default()); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "token", + "@a:example.org", + 0, + )) + .unwrap(); + let failed_session = state.session().unwrap(); + let fetch_state = Arc::clone(&state); + let fetch_url = media_url.clone(); + let fetch = + tokio::spawn(async move { fetch_test(&fetch_state, failed_session, fetch_url).await }); + + assert_eq!(receive_auth(&auth_receiver).await, "Bearer token"); + state.clear_session().unwrap(); + + assert_eq!(fetch.await.unwrap().unwrap_err(), StatusCode::UNAUTHORIZED); + assert!(auth_receiver.try_recv().is_err()); + server.join().unwrap(); + } + + #[tokio::test] + async fn origin_or_scope_switch_while_waiting_terminates_without_a_retry() { + for (origin, scope) in [ + ("https://other.example.org", "@a:example.org"), + ("https://matrix.example.org", "@b:example.org"), + ] { + let state = Arc::new(MediaSessionState::default()); + state + .set_session(test_session( + "https://matrix.example.org", + "old-token", + "@a:example.org", + 0, + )) + .unwrap(); + let previous = state.session().unwrap(); + let waiter_state = Arc::clone(&state); + let waiter = tokio::spawn(async move { + waiter_state + .wait_for_newer_session_with_timeout(&previous, Duration::from_secs(1)) + .await + }); + + state + .set_session(test_session(origin, "new-token", scope, 0)) + .unwrap(); + + assert!(tokio::time::timeout(Duration::from_millis(100), waiter) + .await + .unwrap() + .unwrap() + .is_none()); + } + } + + #[tokio::test] + async fn concurrent_cache_misses_share_the_initial_attempt_and_retry() { + let (media_url, auth_receiver, server) = start_upstream(vec![(401, ""), (200, "data")]); + let state = Arc::new(MediaSessionState::default()); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "old-token", + "@a:example.org", + 0, + )) + .unwrap(); + let failed_session = state.session().unwrap(); + let first_state = Arc::clone(&state); + let first_url = media_url.clone(); + let first = + tokio::spawn(async move { fetch_test(&first_state, failed_session, first_url).await }); + let second_state = Arc::clone(&state); + let second_session = state.session().unwrap(); + let second_url = media_url.clone(); + let second = + tokio::spawn( + async move { fetch_test(&second_state, second_session, second_url).await }, + ); + + assert_eq!(receive_auth(&auth_receiver).await, "Bearer old-token"); + state + .set_session(test_session( + media_url.origin().ascii_serialization().as_str(), + "new-token", + "@a:example.org", + 0, + )) + .unwrap(); + + let first_result = first.await.unwrap(); + let second_result = second.await.unwrap(); + assert!(first_result.is_ok(), "first result: {first_result:?}"); + assert!(second_result.is_ok(), "second result: {second_result:?}"); + assert_eq!(receive_auth(&auth_receiver).await, "Bearer new-token"); + assert!(auth_receiver.try_recv().is_err()); + server.join().unwrap(); + } + + #[test] + fn mismatched_account_or_origin_never_qualifies_for_retry() { + let failed = test_session("https://matrix.example.org", "old", "@a:example.org", 1); + assert!(!should_retry_with_session( + &failed, + &test_session("https://matrix.example.org", "new", "@b:example.org", 2) + )); + assert!(!should_retry_with_session( + &failed, + &test_session("https://other.example.org", "new", "@a:example.org", 2) + )); + } + + #[test] + fn session_marker_is_an_equality_guard_and_validates_decoding() { + let matching: Uri = + "https://sable-media.localhost/media?__sable_media_session=%40a%3Aexample.org" + .parse() + .unwrap(); + let mismatched: Uri = + "https://sable-media.localhost/media?__sable_media_session=%40b%3Aexample.org" + .parse() + .unwrap(); + let malformed: Uri = "https://sable-media.localhost/media?__sable_media_session=%FF" + .parse() + .unwrap(); + let markerless: Uri = "https://sable-media.localhost/media".parse().unwrap(); + + assert!(session_marker_matches(&matching, "@a:example.org").unwrap()); + assert!(!session_marker_matches(&mismatched, "@a:example.org").unwrap()); + assert_eq!( + session_marker_matches(&malformed, "@a:example.org"), + Err(StatusCode::BAD_REQUEST) + ); + assert!(session_marker_matches(&markerless, "@a:example.org").unwrap()); + } + #[tokio::test] async fn does_not_wait_once_a_session_has_been_cleared() { // After logout there is nothing to wait for, so in-flight requests must not hang. diff --git a/src/client/oidcTokenRefresher.test.ts b/src/client/oidcTokenRefresher.test.ts index 39bd1fe9b9..000e9d5b3c 100644 --- a/src/client/oidcTokenRefresher.test.ts +++ b/src/client/oidcTokenRefresher.test.ts @@ -65,6 +65,10 @@ describe('createSessionTokenRefresher', () => { mocks.pushSessionToSW.mockReset().mockResolvedValue(undefined); getAuthMetadata.mockReset().mockResolvedValue(metadata); localStorage.setItem(MATRIX_SESSIONS_KEY, JSON.stringify([session])); + Object.defineProperty(navigator, 'locks', { + configurable: true, + value: undefined, + }); }); it('serializes concurrent same-user refreshers and reuses rotated tokens', async () => { @@ -97,4 +101,31 @@ describe('createSessionTokenRefresher', () => { refreshToken: 'new-refresh', }); }); + + it('uses the current stored token after waiting for the cross-tab refresh lock', async () => { + const request = vi.fn<(name: string, callback: () => Promise) => Promise>( + async (_name, callback) => { + localStorage.setItem( + MATRIX_SESSIONS_KEY, + JSON.stringify([ + { ...session, accessToken: 'other-access', refreshToken: 'other-refresh' }, + ]) + ); + return callback(); + } + ); + Object.defineProperty(navigator, 'locks', { + configurable: true, + value: { request }, + }); + + const refresher = createSessionTokenRefresher(session, mx)!; + + await expect(refresher.tokenRefreshFunction('old-refresh')).resolves.toEqual({ + accessToken: 'other-access', + refreshToken: 'other-refresh', + }); + expect(request).toHaveBeenCalledWith('sable-oidc-refresh', expect.any(Function)); + expect(mocks.refresh).not.toHaveBeenCalled(); + }); }); diff --git a/src/client/oidcTokenRefresher.ts b/src/client/oidcTokenRefresher.ts index 3598adafec..0b32f8c6b2 100644 --- a/src/client/oidcTokenRefresher.ts +++ b/src/client/oidcTokenRefresher.ts @@ -26,6 +26,18 @@ export const assertAuthMetadataIssuer = ( export type SessionTokenRefresher = Pick; const refreshQueues = new Map>(); +const inFlightRefreshes = new Map>(); +const OIDC_REFRESH_LOCK_NAME = 'sable-oidc-refresh'; + +type WebLocks = { + request(name: string, callback: () => Promise): Promise; +}; + +const getWebLocks = (): WebLocks | undefined => { + if (typeof navigator === 'undefined') return undefined; + const locks = (navigator as Navigator & { locks?: WebLocks }).locks; + return locks && typeof locks.request === 'function' ? locks : undefined; +}; const getStoredSession = (userId: string): Session | undefined => getLocalStorageItem(MATRIX_SESSIONS_KEY, []).find( @@ -43,6 +55,10 @@ const withRefreshQueue = async (userId: string, operation: () => Promise): return run; }; +export const waitForSessionTokenRefresh = async (userId: string): Promise => { + await inFlightRefreshes.get(userId)?.catch(() => undefined); +}; + export const createSessionTokenRefresher = ( session: Session, mx: MatrixClient @@ -87,27 +103,42 @@ export const createSessionTokenRefresher = ( return { tokenRefreshFunction: async (refreshToken) => { - return withRefreshQueue(session.userId, async () => { + const refresh = withRefreshQueue(session.userId, async () => { const tokenRefresher = await getTokenRefresher(); - // Another tab may have rotated the token; reusing a consumed one revokes the session. - const storedSession = getStoredSession(session.userId); - const latestRefreshToken = - storedSession?.refreshToken ?? - getStoredSessionRefreshToken(session.userId) ?? - refreshToken; - if (storedSession && latestRefreshToken !== refreshToken) { + const refreshWithCurrentState = async () => { + // Another tab may have rotated the token; reusing a consumed one revokes the session. + const storedSession = getStoredSession(session.userId); + const latestRefreshToken = + storedSession?.refreshToken ?? + getStoredSessionRefreshToken(session.userId) ?? + refreshToken; + if (storedSession && latestRefreshToken !== refreshToken) { + return { + accessToken: storedSession.accessToken, + refreshToken: latestRefreshToken, + }; + } + const tokens = await tokenRefresher.tokenRefreshFunction(latestRefreshToken); return { - accessToken: storedSession.accessToken, - refreshToken: latestRefreshToken, + ...tokens, + // OAuth servers may omit a replacement refresh token, in which case the old one remains valid. + refreshToken: tokens.refreshToken ?? latestRefreshToken, }; - } - const tokens = await tokenRefresher.tokenRefreshFunction(latestRefreshToken); - return { - ...tokens, - // OAuth servers may omit a replacement refresh token, in which case the old one remains valid. - refreshToken: tokens.refreshToken ?? latestRefreshToken, }; + + const locks = getWebLocks(); + return locks + ? locks.request(OIDC_REFRESH_LOCK_NAME, refreshWithCurrentState) + : refreshWithCurrentState(); }); + inFlightRefreshes.set(session.userId, refresh); + try { + return await refresh; + } finally { + if (inFlightRefreshes.get(session.userId) === refresh) { + inFlightRefreshes.delete(session.userId); + } + } }, }; }; diff --git a/src/serviceWorkerBootstrap.test.ts b/src/serviceWorkerBootstrap.test.ts index aa3c866526..01b1ba8039 100644 --- a/src/serviceWorkerBootstrap.test.ts +++ b/src/serviceWorkerBootstrap.test.ts @@ -7,6 +7,8 @@ const { mockAddEventListener, mockReady, mockPushSessionToSW, + mockWaitForSessionTokenRefresh, + mockGetLocalStorageItem, mockWarn, } = vi.hoisted(() => ({ mockHasServiceWorker: vi.fn<() => boolean>(), @@ -32,6 +34,8 @@ const { >(), mockReady: Promise.resolve(undefined), mockPushSessionToSW: vi.fn<(baseUrl?: string, accessToken?: string, userId?: string) => void>(), + mockWaitForSessionTokenRefresh: vi.fn<(userId: string) => Promise>(), + mockGetLocalStorageItem: vi.fn<(key: string, fallback: unknown) => unknown>(), mockWarn: vi.fn<(...args: unknown[]) => void>(), })); @@ -44,6 +48,10 @@ vi.mock('./sw-session', () => ({ pushSessionToSW: mockPushSessionToSW, })); +vi.mock('./client/oidcTokenRefresher', () => ({ + waitForSessionTokenRefresh: mockWaitForSessionTokenRefresh, +})); + vi.mock('./app/state/sessions', () => ({ getFallbackSession: vi.fn<() => undefined>(() => undefined), MATRIX_SESSIONS_KEY: 'matrix-sessions', @@ -51,9 +59,7 @@ vi.mock('./app/state/sessions', () => ({ })); vi.mock('./app/state/utils/atomWithLocalStorage', () => ({ - getLocalStorageItem: vi.fn<(key: string, fallback: unknown) => unknown>( - (_: string, fallback: unknown) => fallback - ), + getLocalStorageItem: mockGetLocalStorageItem, })); vi.mock('./app/utils/debug', () => ({ @@ -66,6 +72,8 @@ describe('registerAppServiceWorker', () => { beforeEach(() => { vi.clearAllMocks(); mockHasServiceWorker.mockReturnValue(false); + mockWaitForSessionTokenRefresh.mockReset().mockResolvedValue(undefined); + mockGetLocalStorageItem.mockImplementation((_key, fallback) => fallback); Object.defineProperty(window, 'confirm', { configurable: true, value: vi.fn<(message?: string) => boolean>(() => false), @@ -128,4 +136,45 @@ describe('registerAppServiceWorker', () => { expect(mockPushSessionToSW).toHaveBeenCalledTimes(1); }); + + it('waits for the active session refresh before replying to a session request', async () => { + let activeSession: { baseUrl: string; userId: string; accessToken: string } | undefined; + mockGetLocalStorageItem.mockImplementation((key, fallback) => { + if (key === 'matrix-sessions') return activeSession ? [activeSession] : []; + if (key === 'active-session') return activeSession?.userId; + return fallback; + }); + mockHasServiceWorker.mockReturnValue(true); + registerAppServiceWorker(); + await Promise.resolve(); + await Promise.resolve(); + mockPushSessionToSW.mockReset(); + + activeSession = { + baseUrl: 'https://hs.example', + userId: '@alice:hs.example', + accessToken: 'new-access', + }; + let resolveRefresh!: () => void; + mockWaitForSessionTokenRefresh.mockReturnValueOnce( + new Promise((resolve) => { + resolveRefresh = resolve; + }) + ); + const listener = mockAddEventListener.mock.calls.find(([type]) => type === 'message')?.[1] as + | ((event: MessageEvent) => void) + | undefined; + listener?.({ data: { type: 'requestSession' } } as MessageEvent); + + expect(mockWaitForSessionTokenRefresh).toHaveBeenCalledWith('@alice:hs.example'); + expect(mockPushSessionToSW).not.toHaveBeenCalled(); + resolveRefresh(); + await vi.waitFor(() => + expect(mockPushSessionToSW).toHaveBeenCalledWith( + 'https://hs.example', + 'new-access', + '@alice:hs.example' + ) + ); + }); }); diff --git a/src/serviceWorkerBootstrap.ts b/src/serviceWorkerBootstrap.ts index 5b515f9bfc..4fcb3d953c 100644 --- a/src/serviceWorkerBootstrap.ts +++ b/src/serviceWorkerBootstrap.ts @@ -5,8 +5,25 @@ import { getFallbackSession, MATRIX_SESSIONS_KEY, ACTIVE_SESSION_KEY } from './a import { getLocalStorageItem } from './app/state/utils/atomWithLocalStorage'; import { hasServiceWorker } from './app/utils/platform'; import { pushSessionToSW } from './sw-session'; +import { waitForSessionTokenRefresh } from './client/oidcTokenRefresher'; const log = createLogger('service-worker-bootstrap'); +const REFRESH_WAIT_TIMEOUT_MS = 2500; + +const waitForRefreshWithTimeout = async (userId: string): Promise => { + let timeoutId: ReturnType | undefined; + const timeout = new Promise((resolve) => { + timeoutId = setTimeout(resolve, REFRESH_WAIT_TIMEOUT_MS); + }); + await Promise.race([waitForSessionTokenRefresh(userId), timeout]); + if (timeoutId !== undefined) clearTimeout(timeoutId); +}; + +const getActiveSession = () => { + const sessions = getLocalStorageItem(MATRIX_SESSIONS_KEY, []); + const activeId = getLocalStorageItem(ACTIVE_SESSION_KEY, undefined); + return sessions.find((s) => s.userId === activeId) ?? sessions[0] ?? getFallbackSession(); +}; const showUpdateAvailablePrompt = (registration: ServiceWorkerRegistration) => { const DONT_SHOW_PROMPT_KEY = 'cinny_dont_show_sw_update_prompt'; @@ -23,11 +40,11 @@ const showUpdateAvailablePrompt = (registration: ServiceWorkerRegistration) => { ); }; -const sendSessionToSW = () => { - const sessions = getLocalStorageItem(MATRIX_SESSIONS_KEY, []); - const activeId = getLocalStorageItem(ACTIVE_SESSION_KEY, undefined); - const active = sessions.find((s) => s.userId === activeId) ?? sessions[0] ?? getFallbackSession(); - pushSessionToSW(active?.baseUrl, active?.accessToken, active?.userId); +const sendSessionToSW = async () => { + const active = getActiveSession(); + if (active) await waitForRefreshWithTimeout(active.userId); + const current = getActiveSession(); + await pushSessionToSW(current?.baseUrl, current?.accessToken, current?.userId); }; export function registerAppServiceWorker() { @@ -76,7 +93,7 @@ export function registerAppServiceWorker() { const { type } = data as { type?: unknown }; if (type === 'requestSession') { - sendSessionToSW(); + void sendSessionToSW(); } if (data.type === 'token' && data.id) { diff --git a/src/sw-media-auth-recovery.test.ts b/src/sw-media-auth-recovery.test.ts index d3ef26fe88..ee96d8f68c 100644 --- a/src/sw-media-auth-recovery.test.ts +++ b/src/sw-media-auth-recovery.test.ts @@ -127,4 +127,32 @@ describe('service worker media auth recovery', () => { swTestHooks.setSession(client.id, 'new-token', 'https://matrix.example.org'); await expect(nextRequest).resolves.toMatchObject({ accessToken: 'new-token' }); }); + + it('does not retry when the refreshed session has the same access token', async () => { + const client = { + id: 'client-same-token', + postMessage: vi.fn(), + } as unknown as Client; + clients.set(client.id, client); + const session = { accessToken: 'same-token', baseUrl: 'https://matrix.example.org' }; + vi.mocked(fetch).mockResolvedValue( + new Response(JSON.stringify({ errcode: 'M_UNKNOWN_TOKEN' }), { status: 401 }) + ); + + const request = new Request( + 'https://matrix.example.org/_matrix/client/v1/media/download/example.org/media-id' + ); + const recovery = swTestHooks.respondWithMediaAuthRecovery( + request, + session, + 'follow', + client.id + ); + + await vi.waitFor(() => expect(client.postMessage).toHaveBeenCalledTimes(1)); + swTestHooks.setSession(client.id, 'same-token', session.baseUrl); + + await expect(recovery).resolves.toHaveProperty('status', 401); + expect(fetch).toHaveBeenCalledTimes(1); + }); });