diff --git a/crates/matrix-sdk-base/Cargo.toml b/crates/matrix-sdk-base/Cargo.toml index fdec9ee28..e160b2b3b 100644 --- a/crates/matrix-sdk-base/Cargo.toml +++ b/crates/matrix-sdk-base/Cargo.toml @@ -37,6 +37,7 @@ http = { version = "0.2.6", optional = true } lru = "0.7.5" matrix-sdk-common = { version = "0.4.0", path = "../matrix-sdk-common" } matrix-sdk-crypto = { version = "0.4.0", path = "../matrix-sdk-crypto", optional = true } +once_cell = "1.10.0" pbkdf2 = { version = "0.11.0", default-features = false, optional = true } rand = { version = "0.8.5", optional = true } serde = { version = "1.0.136", features = ["rc"] } diff --git a/crates/matrix-sdk-base/src/client.rs b/crates/matrix-sdk-base/src/client.rs index 9008bd7fb..b32a4236e 100644 --- a/crates/matrix-sdk-base/src/client.rs +++ b/crates/matrix-sdk-base/src/client.rs @@ -86,9 +86,6 @@ pub type Token = String; /// accordingly updates its state. #[derive(Clone)] pub struct BaseClient { - /// The current client session containing our user id, device id and access - /// token. - session: Arc>>, /// The current sync token that should be used for the next sync call. pub(crate) sync_token: Arc>>, /// Database @@ -101,7 +98,7 @@ pub struct BaseClient { impl fmt::Debug for BaseClient { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { f.debug_struct("Client") - .field("session", &self.session) + .field("session", &self.session()) .field("sync_token", &self.sync_token) .finish() } @@ -169,7 +166,6 @@ impl BaseClient { let holder = config.crypto_store.map(CryptoHolder::new).unwrap_or_default(); BaseClient { - session: store.session.clone(), sync_token: store.sync_token.clone(), store, #[cfg(feature = "e2e-encryption")] @@ -177,10 +173,16 @@ impl BaseClient { } } - /// The current client session containing our user id, device id and access - /// token. - pub fn session(&self) -> &Arc>> { - &self.session + /// Get the user login session. + /// + /// If the client is currently logged in, this will return a + /// [`matrix_sdk::Session`] object which can later be given to + /// `restore_login`. + /// + /// Returns a session object if the client is logged in. Otherwise returns + /// `None`. + pub fn session(&self) -> Option<&Session> { + self.store.session() } /// Get a reference to the store. @@ -189,10 +191,8 @@ impl BaseClient { } /// Is the client logged in. - pub async fn logged_in(&self) -> bool { - // TODO turn this into a atomic bool so this method doesn't need to be - // async. - self.session.read().await.is_some() + pub fn logged_in(&self) -> bool { + self.store.session().is_some() } /// Receive a login response and update the session of the client. @@ -229,8 +229,6 @@ impl BaseClient { olm.convert_to_olm(&session).await?; } - *self.session.write().await = Some(session); - Ok(()) } @@ -1127,18 +1125,6 @@ impl BaseClient { } } - /// Get the user login session. - /// - /// If the client is currently logged in, this will return a - /// `matrix_sdk::Session` object which can later be given to - /// `restore_login`. - /// - /// Returns a session object if the client is logged in. Otherwise returns - /// `None`. - pub async fn get_session(&self) -> Option { - self.session.read().await.clone() - } - /// Get a map holding all the devices of an user. /// /// This will always return an empty map if the client hasn't been logged @@ -1210,7 +1196,7 @@ impl BaseClient { .transpose()? { Ok(event.content.global) - } else if let Some(session) = self.get_session().await { + } else if let Some(session) = self.session() { Ok(Ruleset::server_default(&session.user_id)) } else { Ok(Ruleset::new()) diff --git a/crates/matrix-sdk-base/src/lib.rs b/crates/matrix-sdk-base/src/lib.rs index f32253e60..3924d80b2 100644 --- a/crates/matrix-sdk-base/src/lib.rs +++ b/crates/matrix-sdk-base/src/lib.rs @@ -41,6 +41,7 @@ pub use client::BaseClient; pub use http; #[cfg(feature = "e2e-encryption")] pub use matrix_sdk_crypto as crypto; +pub use once_cell; pub use rooms::{DisplayName, Room, RoomInfo, RoomMember, RoomType}; pub use store::{StateChanges, StateStore, Store, StoreError}; pub use utils::{ diff --git a/crates/matrix-sdk-base/src/store/mod.rs b/crates/matrix-sdk-base/src/store/mod.rs index 3d4c99cdd..b7a581cf2 100644 --- a/crates/matrix-sdk-base/src/store/mod.rs +++ b/crates/matrix-sdk-base/src/store/mod.rs @@ -28,6 +28,8 @@ use std::{ sync::Arc, }; +use once_cell::sync::OnceCell; + #[cfg(any(test, feature = "testing"))] #[macro_use] pub mod integration_tests; @@ -380,7 +382,7 @@ pub trait StateStore: AsyncTraitDeps { #[derive(Debug, Clone)] pub struct Store { inner: Arc, - pub(crate) session: Arc>>, + pub(crate) session: Arc>, pub(crate) sync_token: Arc>>, rooms: Arc>, stripped_rooms: Arc>, @@ -423,11 +425,17 @@ impl Store { let token = self.get_sync_token().await?; *self.sync_token.write().await = token; - *self.session.write().await = Some(session); + self.session.set(session).expect("A session was already set"); Ok(()) } + /// The current [`Session`] containing our user id, device id and access + /// token. + pub fn session(&self) -> Option<&Session> { + self.session.get() + } + /// Get all the rooms this store knows about. pub fn get_rooms(&self) -> Vec { self.rooms.iter().filter_map(|r| self.get_room(r.key())).collect() @@ -458,8 +466,7 @@ impl Store { /// Lookup the stripped Room for the given RoomId, or create one, if it /// didn't exist yet in the store pub async fn get_or_create_stripped_room(&self, room_id: &RoomId) -> Room { - let session = self.session.read().await; - let user_id = &session.as_ref().expect("Creating room while not being logged in").user_id; + let user_id = &self.session().expect("Creating room while not being logged in").user_id; self.stripped_rooms .entry(room_id.to_owned()) @@ -474,8 +481,7 @@ impl Store { return self.get_or_create_stripped_room(room_id).await; } - let session = self.session.read().await; - let user_id = &session.as_ref().expect("Creating room while not being logged in").user_id; + let user_id = &self.session().expect("Creating room while not being logged in").user_id; self.rooms .entry(room_id.to_owned()) diff --git a/crates/matrix-sdk/src/account.rs b/crates/matrix-sdk/src/account.rs index 995095e5c..53451c565 100644 --- a/crates/matrix-sdk/src/account.rs +++ b/crates/matrix-sdk/src/account.rs @@ -70,8 +70,8 @@ impl Account { /// # Result::<_, matrix_sdk::Error>::Ok(()) }); /// ``` pub async fn get_display_name(&self) -> Result> { - let user_id = self.client.user_id().await.ok_or(Error::AuthenticationRequired)?; - let request = get_display_name::v3::Request::new(&user_id); + let user_id = self.client.user_id().ok_or(Error::AuthenticationRequired)?; + let request = get_display_name::v3::Request::new(user_id); let response = self.client.send(request, None).await?; Ok(response.displayname) } @@ -93,8 +93,8 @@ impl Account { /// # Result::<_, matrix_sdk::Error>::Ok(()) }); /// ``` pub async fn set_display_name(&self, name: Option<&str>) -> Result<()> { - let user_id = self.client.user_id().await.ok_or(Error::AuthenticationRequired)?; - let request = set_display_name::v3::Request::new(&user_id, name); + let user_id = self.client.user_id().ok_or(Error::AuthenticationRequired)?; + let request = set_display_name::v3::Request::new(user_id, name); self.client.send(request, None).await?; Ok(()) } @@ -118,8 +118,8 @@ impl Account { /// # Result::<_, matrix_sdk::Error>::Ok(()) }); /// ``` pub async fn get_avatar_url(&self) -> Result> { - let user_id = self.client.user_id().await.ok_or(Error::AuthenticationRequired)?; - let request = get_avatar_url::v3::Request::new(&user_id); + let user_id = self.client.user_id().ok_or(Error::AuthenticationRequired)?; + let request = get_avatar_url::v3::Request::new(user_id); let config = Some(RequestConfig::new().force_auth()); @@ -131,8 +131,8 @@ impl Account { /// /// The avatar is unset if `url` is `None`. pub async fn set_avatar_url(&self, url: Option<&MxcUri>) -> Result<()> { - let user_id = self.client.user_id().await.ok_or(Error::AuthenticationRequired)?; - let request = set_avatar_url::v3::Request::new(&user_id, url); + let user_id = self.client.user_id().ok_or(Error::AuthenticationRequired)?; + let request = set_avatar_url::v3::Request::new(user_id, url); self.client.send(request, None).await?; Ok(()) } @@ -233,8 +233,8 @@ impl Account { /// # Result::<_, matrix_sdk::Error>::Ok(()) }); /// ``` pub async fn get_profile(&self) -> Result { - let user_id = self.client.user_id().await.ok_or(Error::AuthenticationRequired)?; - let request = get_profile::v3::Request::new(&user_id); + let user_id = self.client.user_id().ok_or(Error::AuthenticationRequired)?; + let request = get_profile::v3::Request::new(user_id); Ok(self.client.send(request, None).await?) } diff --git a/crates/matrix-sdk/src/client/builder.rs b/crates/matrix-sdk/src/client/builder.rs index bfa6440bd..3ae63d070 100644 --- a/crates/matrix-sdk/src/client/builder.rs +++ b/crates/matrix-sdk/src/client/builder.rs @@ -291,12 +291,7 @@ impl ClientBuilder { let base_client = BaseClient::with_store_config(self.store_config); let mk_http_client = |homeserver| { - HttpClient::new( - inner_http_client.clone(), - homeserver, - base_client.session().clone(), - self.request_config, - ) + HttpClient::new(inner_http_client.clone(), homeserver, self.request_config) }; let homeserver = match homeserver_cfg { diff --git a/crates/matrix-sdk/src/client/mod.rs b/crates/matrix-sdk/src/client/mod.rs index 06125ad5d..14a9ea88c 100644 --- a/crates/matrix-sdk/src/client/mod.rs +++ b/crates/matrix-sdk/src/client/mod.rs @@ -64,8 +64,8 @@ use ruma::{ assign, events::room::MediaSource, presence::PresenceState, - MxcUri, OwnedDeviceId, OwnedRoomId, OwnedServerName, OwnedUserId, RoomId, RoomOrAliasId, - ServerName, UInt, + DeviceId, MxcUri, OwnedDeviceId, OwnedRoomId, OwnedServerName, RoomId, RoomOrAliasId, + ServerName, UInt, UserId, }; use serde::de::DeserializeOwned; #[cfg(not(target_arch = "wasm32"))] @@ -266,8 +266,8 @@ impl Client { } /// Is the client logged in. - pub async fn logged_in(&self) -> bool { - self.inner.base_client.logged_in().await + pub fn logged_in(&self) -> bool { + self.inner.base_client.logged_in() } /// The Homeserver of the client. @@ -276,15 +276,13 @@ impl Client { } /// Get the user id of the current owner of the client. - pub async fn user_id(&self) -> Option { - let session = self.inner.base_client.session().read().await; - session.as_ref().cloned().map(|s| s.user_id) + pub fn user_id(&self) -> Option<&UserId> { + self.inner.base_client.session().map(|s| s.user_id.as_ref()) } /// Get the device id that identifies the current session. - pub async fn device_id(&self) -> Option { - let session = self.inner.base_client.session().read().await; - session.as_ref().map(|s| s.device_id.clone()) + pub fn device_id(&self) -> Option<&DeviceId> { + self.inner.base_client.session().map(|s| s.device_id.as_ref()) } /// Get the whole session info of this client. @@ -293,8 +291,8 @@ impl Client { /// /// Can be used with [`Client::restore_login`] to restore a previously /// logged in session. - pub async fn session(&self) -> Option { - self.inner.base_client.session().read().await.clone() + pub fn session(&self) -> Option<&Session> { + self.inner.base_client.session() } /// Get a reference to the store. @@ -1099,6 +1097,7 @@ impl Client { /// /// [`login`]: #method.login pub async fn restore_login(&self, session: Session) -> Result<()> { + self.inner.http_client.set_session(session.clone()); Ok(self.inner.base_client.restore_login(session).await?) } @@ -1210,8 +1209,8 @@ impl Client { if let Some(filter) = self.inner.base_client.get_filter(filter_name).await? { Ok(filter) } else { - let user_id = self.user_id().await.ok_or(Error::AuthenticationRequired)?; - let request = FilterUploadRequest::new(&user_id, definition); + let user_id = self.user_id().ok_or(Error::AuthenticationRequired)?; + let request = FilterUploadRequest::new(user_id, definition); let response = self.send(request, None).await?; self.inner.base_client.receive_filter_upload(filter_name, &response).await?; diff --git a/crates/matrix-sdk/src/encryption/mod.rs b/crates/matrix-sdk/src/encryption/mod.rs index 6b50c2a49..ba5220dc2 100644 --- a/crates/matrix-sdk/src/encryption/mod.rs +++ b/crates/matrix-sdk/src/encryption/mod.rs @@ -245,9 +245,9 @@ impl Client { T: GlobalAccountDataEventContent, { let own_user = - self.user_id().await.ok_or_else(|| Error::from(HttpError::AuthenticationRequired))?; + self.user_id().ok_or_else(|| Error::from(HttpError::AuthenticationRequired))?; - let request = set_global_account_data::v3::Request::new(&content, &own_user)?; + let request = set_global_account_data::v3::Request::new(&content, own_user)?; Ok(self.send(request, None).await?) } diff --git a/crates/matrix-sdk/src/http_client.rs b/crates/matrix-sdk/src/http_client.rs index 9161e76a1..1ebb5f502 100644 --- a/crates/matrix-sdk/src/http_client.rs +++ b/crates/matrix-sdk/src/http_client.rs @@ -17,6 +17,7 @@ use std::{any::type_name, convert::TryFrom, fmt::Debug, sync::Arc, time::Duratio use async_trait::async_trait; use bytes::{Bytes, BytesMut}; use http::Response as HttpResponse; +use matrix_sdk_base::once_cell::sync::OnceCell; use matrix_sdk_common::{locks::RwLock, AsyncTraitDeps}; use reqwest::Response; use ruma::api::{ @@ -95,7 +96,7 @@ pub trait HttpSend: AsyncTraitDeps { pub(crate) struct HttpClient { pub(crate) inner: Arc, pub(crate) homeserver: Arc>, - pub(crate) session: Arc>>, + pub(crate) session: OnceCell, pub(crate) request_config: RequestConfig, } @@ -103,10 +104,9 @@ impl HttpClient { pub(crate) fn new( inner: Arc, homeserver: Arc>, - session: Arc>>, request_config: RequestConfig, ) -> Self { - HttpClient { inner, homeserver, session, request_config } + HttpClient { inner, homeserver, session: Default::default(), request_config } } #[tracing::instrument(skip(self, request), fields(request_type = type_name::()))] @@ -130,21 +130,18 @@ impl HttpClient { return Err(HttpError::NotClientRequest); } - let access_token; - let request = if !self.request_config.assert_identity { let send_access_token = if auth_scheme == AuthScheme::None && !config.force_auth { // Small optimization: Don't take the session lock if we know the auth token // isn't going to be used anyways. SendAccessToken::None } else { - match self.session.read().await.as_ref() { + match self.session() { Some(session) => { - access_token = session.access_token.clone(); if config.force_auth { - SendAccessToken::Always(&access_token) + SendAccessToken::Always(&session.access_token) } else { - SendAccessToken::IfRequired(&access_token) + SendAccessToken::IfRequired(&session.access_token) } } None => SendAccessToken::None, @@ -157,18 +154,12 @@ impl HttpClient { &server_versions, )? } else { - let (send_access_token, user_id) = { - let session = self.session.read().await; - let session = session.as_ref().ok_or(HttpError::UserIdRequired)?; - - access_token = session.access_token.clone(); - (SendAccessToken::Always(&access_token), session.user_id.clone()) - }; - request.try_into_http_request_with_user_id::( &self.homeserver.read().await.to_string(), - send_access_token, - &user_id, + SendAccessToken::Always( + &self.session().ok_or(HttpError::UserIdRequired)?.access_token, + ), + &self.session().ok_or(HttpError::UserIdRequired)?.user_id, &server_versions, )? }; @@ -182,6 +173,14 @@ impl HttpClient { Ok(response) } + + pub(crate) fn set_session(&self, session: Session) { + self.session.set(session).expect("A session was already set"); + } + + fn session(&self) -> Option<&Session> { + self.session.get() + } } #[derive(Debug)] diff --git a/crates/matrix-sdk/src/room/common.rs b/crates/matrix-sdk/src/room/common.rs index b6ea52a4c..433102877 100644 --- a/crates/matrix-sdk/src/room/common.rs +++ b/crates/matrix-sdk/src/room/common.rs @@ -809,9 +809,9 @@ impl Common { tag: TagName, tag_info: TagInfo, ) -> HttpResult { - let user_id = self.client.user_id().await.ok_or(HttpError::AuthenticationRequired)?; + let user_id = self.client.user_id().ok_or(HttpError::AuthenticationRequired)?; let request = - create_tag::v3::Request::new(&user_id, self.inner.room_id(), tag.as_ref(), tag_info); + create_tag::v3::Request::new(user_id, self.inner.room_id(), tag.as_ref(), tag_info); self.client.send(request, None).await } @@ -822,8 +822,8 @@ impl Common { /// # Arguments /// * `tag` - The tag to remove. pub async fn remove_tag(&self, tag: TagName) -> HttpResult { - let user_id = self.client.user_id().await.ok_or(HttpError::AuthenticationRequired)?; - let request = delete_tag::v3::Request::new(&user_id, self.inner.room_id(), tag.as_ref()); + let user_id = self.client.user_id().ok_or(HttpError::AuthenticationRequired)?; + let request = delete_tag::v3::Request::new(user_id, self.inner.room_id(), tag.as_ref()); self.client.send(request, None).await } @@ -836,11 +836,8 @@ impl Common { /// # Arguments /// * `is_direct` - Whether to mark this room as direct. pub async fn set_is_direct(&self, is_direct: bool) -> Result<()> { - let user_id = self - .client - .user_id() - .await - .ok_or_else(|| Error::from(HttpError::AuthenticationRequired))?; + let user_id = + self.client.user_id().ok_or_else(|| Error::from(HttpError::AuthenticationRequired))?; let mut content = self .client @@ -871,7 +868,7 @@ impl Common { content.retain(|_, list| !list.is_empty()); } - let request = set_global_account_data::v3::Request::new(&content, &user_id)?; + let request = set_global_account_data::v3::Request::new(&content, user_id)?; self.client.send(request, None).await?; Ok(())