Merge pull request #3299 from element-hq/feature/bma/rotateSessionPath

Ensure sessionPath is not reused for different homeserver. Fixes not loading media issue.
This commit is contained in:
Benoit Marty 2024-08-12 16:31:46 +02:00 committed by GitHub
commit 2b7f78255f

View file

@ -60,7 +60,7 @@ import javax.inject.Inject
@ContributesBinding(AppScope::class) @ContributesBinding(AppScope::class)
@SingleIn(AppScope::class) @SingleIn(AppScope::class)
class RustMatrixAuthenticationService @Inject constructor( class RustMatrixAuthenticationService @Inject constructor(
baseDirectory: File, private val baseDirectory: File,
private val coroutineDispatchers: CoroutineDispatchers, private val coroutineDispatchers: CoroutineDispatchers,
private val sessionStore: SessionStore, private val sessionStore: SessionStore,
private val rustMatrixClientFactory: RustMatrixClientFactory, private val rustMatrixClientFactory: RustMatrixClientFactory,
@ -70,10 +70,19 @@ class RustMatrixAuthenticationService @Inject constructor(
// Passphrase which will be used for new sessions. Existing sessions will use the passphrase // Passphrase which will be used for new sessions. Existing sessions will use the passphrase
// stored in the SessionData. // stored in the SessionData.
private val pendingPassphrase = getDatabasePassphrase() private val pendingPassphrase = getDatabasePassphrase()
private val sessionPath = File(baseDirectory, UUID.randomUUID().toString()).absolutePath
// Need to keep a copy of the current session path to eventually delete it.
// Ideally it would be possible to get the sessionPath from the Client to avoid doing this.
private var sessionPath: File? = null
private var currentClient: Client? = null private var currentClient: Client? = null
private var currentHomeserver = MutableStateFlow<MatrixHomeServerDetails?>(null) private var currentHomeserver = MutableStateFlow<MatrixHomeServerDetails?>(null)
private fun rotateSessionPath(): File {
sessionPath?.deleteRecursively()
return File(baseDirectory, UUID.randomUUID().toString())
.also { sessionPath = it }
}
override fun loggedInStateFlow(): Flow<LoggedInState> { override fun loggedInStateFlow(): Flow<LoggedInState> {
return sessionStore.isLoggedIn() return sessionStore.isLoggedIn()
} }
@ -117,8 +126,9 @@ class RustMatrixAuthenticationService @Inject constructor(
override suspend fun setHomeserver(homeserver: String): Result<Unit> = override suspend fun setHomeserver(homeserver: String): Result<Unit> =
withContext(coroutineDispatchers.io) { withContext(coroutineDispatchers.io) {
val emptySessionPath = rotateSessionPath()
runCatching { runCatching {
val client = getBaseClientBuilder() val client = getBaseClientBuilder(emptySessionPath)
.serverNameOrHomeserverUrl(homeserver) .serverNameOrHomeserverUrl(homeserver)
.build() .build()
currentClient = client currentClient = client
@ -135,13 +145,14 @@ class RustMatrixAuthenticationService @Inject constructor(
withContext(coroutineDispatchers.io) { withContext(coroutineDispatchers.io) {
runCatching { runCatching {
val client = currentClient ?: error("You need to call `setHomeserver()` first") val client = currentClient ?: error("You need to call `setHomeserver()` first")
val currentSessionPath = sessionPath ?: error("You need to call `setHomeserver()` first")
client.login(username, password, "Element X Android", null) client.login(username, password, "Element X Android", null)
val sessionData = client.session() val sessionData = client.session()
.toSessionData( .toSessionData(
isTokenValid = true, isTokenValid = true,
loginType = LoginType.PASSWORD, loginType = LoginType.PASSWORD,
passphrase = pendingPassphrase, passphrase = pendingPassphrase,
sessionPath = sessionPath, sessionPath = currentSessionPath.absolutePath,
) )
clear() clear()
sessionStore.storeData(sessionData) sessionStore.storeData(sessionData)
@ -185,13 +196,14 @@ class RustMatrixAuthenticationService @Inject constructor(
return withContext(coroutineDispatchers.io) { return withContext(coroutineDispatchers.io) {
runCatching { runCatching {
val client = currentClient ?: error("You need to call `setHomeserver()` first") val client = currentClient ?: error("You need to call `setHomeserver()` first")
val currentSessionPath = sessionPath ?: error("You need to call `setHomeserver()` first")
val urlForOidcLogin = pendingOidcAuthorizationData ?: error("You need to call `getOidcUrl()` first") val urlForOidcLogin = pendingOidcAuthorizationData ?: error("You need to call `getOidcUrl()` first")
client.loginWithOidcCallback(urlForOidcLogin, callbackUrl) client.loginWithOidcCallback(urlForOidcLogin, callbackUrl)
val sessionData = client.session().toSessionData( val sessionData = client.session().toSessionData(
isTokenValid = true, isTokenValid = true,
loginType = LoginType.OIDC, loginType = LoginType.OIDC,
passphrase = pendingPassphrase, passphrase = pendingPassphrase,
sessionPath = sessionPath, sessionPath = currentSessionPath.absolutePath,
) )
clear() clear()
pendingOidcAuthorizationData?.close() pendingOidcAuthorizationData?.close()
@ -206,9 +218,10 @@ class RustMatrixAuthenticationService @Inject constructor(
override suspend fun loginWithQrCode(qrCodeData: MatrixQrCodeLoginData, progress: (QrCodeLoginStep) -> Unit) = override suspend fun loginWithQrCode(qrCodeData: MatrixQrCodeLoginData, progress: (QrCodeLoginStep) -> Unit) =
withContext(coroutineDispatchers.io) { withContext(coroutineDispatchers.io) {
val emptySessionPath = rotateSessionPath()
runCatching { runCatching {
val client = rustMatrixClientFactory.getBaseClientBuilder( val client = rustMatrixClientFactory.getBaseClientBuilder(
sessionPath = sessionPath, sessionPath = emptySessionPath.absolutePath,
passphrase = pendingPassphrase, passphrase = pendingPassphrase,
slidingSyncProxy = AuthenticationConfig.SLIDING_SYNC_PROXY_URL, slidingSyncProxy = AuthenticationConfig.SLIDING_SYNC_PROXY_URL,
slidingSync = ClientBuilderSlidingSync.Discovered, slidingSync = ClientBuilderSlidingSync.Discovered,
@ -229,7 +242,7 @@ class RustMatrixAuthenticationService @Inject constructor(
isTokenValid = true, isTokenValid = true,
loginType = LoginType.QR, loginType = LoginType.QR,
passphrase = pendingPassphrase, passphrase = pendingPassphrase,
sessionPath = sessionPath, sessionPath = emptySessionPath.absolutePath,
) )
sessionStore.storeData(sessionData) sessionStore.storeData(sessionData)
SessionId(sessionData.userId) SessionId(sessionData.userId)
@ -246,11 +259,13 @@ class RustMatrixAuthenticationService @Inject constructor(
} }
Timber.e(throwable, "Failed to login with QR code") Timber.e(throwable, "Failed to login with QR code")
} }
} }
private fun getBaseClientBuilder() = rustMatrixClientFactory private fun getBaseClientBuilder(
sessionPath: File,
) = rustMatrixClientFactory
.getBaseClientBuilder( .getBaseClientBuilder(
sessionPath = sessionPath, sessionPath = sessionPath.absolutePath,
passphrase = pendingPassphrase, passphrase = pendingPassphrase,
slidingSyncProxy = AuthenticationConfig.SLIDING_SYNC_PROXY_URL, slidingSyncProxy = AuthenticationConfig.SLIDING_SYNC_PROXY_URL,
slidingSync = ClientBuilderSlidingSync.Discovered, slidingSync = ClientBuilderSlidingSync.Discovered,