Merge pull request #4400 from element-hq/feature/bma/customPushGateway

Let EnterpriseService provides push gateways
This commit is contained in:
Benoit Marty 2025-03-13 12:01:10 +01:00 committed by GitHub
commit c53d80e8fb
16 changed files with 130 additions and 15 deletions

@ -1 +1 @@
Subproject commit 6d96bf58aec2ecc77b408858272cd64ec26e10d0 Subproject commit 665a15a1907a116816ffd1653bccfeaeeb7a2968

View file

@ -17,4 +17,7 @@ interface EnterpriseService {
fun semanticColorsLight(): SemanticColors fun semanticColorsLight(): SemanticColors
fun semanticColorsDark(): SemanticColors fun semanticColorsDark(): SemanticColors
fun firebasePushGateway(): String?
fun unifiedPushDefaultPushGateway(): String?
} }

View file

@ -27,4 +27,7 @@ class DefaultEnterpriseService @Inject constructor() : EnterpriseService {
override fun semanticColorsLight(): SemanticColors = compoundColorsLight override fun semanticColorsLight(): SemanticColors = compoundColorsLight
override fun semanticColorsDark(): SemanticColors = compoundColorsDark override fun semanticColorsDark(): SemanticColors = compoundColorsDark
override fun firebasePushGateway(): String? = null
override fun unifiedPushDefaultPushGateway(): String? = null
} }

View file

@ -19,6 +19,8 @@ class FakeEnterpriseService(
private val defaultHomeserverResult: () -> String? = { A_FAKE_HOMESERVER }, private val defaultHomeserverResult: () -> String? = { A_FAKE_HOMESERVER },
private val semanticColorsLightResult: () -> SemanticColors = { lambdaError() }, private val semanticColorsLightResult: () -> SemanticColors = { lambdaError() },
private val semanticColorsDarkResult: () -> SemanticColors = { lambdaError() }, private val semanticColorsDarkResult: () -> SemanticColors = { lambdaError() },
private val firebasePushGatewayResult: () -> String? = { lambdaError() },
private val unifiedPushDefaultPushGatewayResult: () -> String? = { lambdaError() },
) : EnterpriseService { ) : EnterpriseService {
override suspend fun isEnterpriseUser(sessionId: SessionId): Boolean = simulateLongTask { override suspend fun isEnterpriseUser(sessionId: SessionId): Boolean = simulateLongTask {
isEnterpriseUserResult(sessionId) isEnterpriseUserResult(sessionId)
@ -36,6 +38,14 @@ class FakeEnterpriseService(
return semanticColorsDarkResult() return semanticColorsDarkResult()
} }
override fun firebasePushGateway(): String? {
return firebasePushGatewayResult()
}
override fun unifiedPushDefaultPushGateway(): String? {
return unifiedPushDefaultPushGatewayResult()
}
companion object { companion object {
const val A_FAKE_HOMESERVER = "a_fake_homeserver" const val A_FAKE_HOMESERVER = "a_fake_homeserver"
} }

View file

@ -50,6 +50,7 @@ setupAnvil()
dependencies { dependencies {
implementation(libs.dagger) implementation(libs.dagger)
implementation(libs.androidx.corektx) implementation(libs.androidx.corektx)
implementation(projects.features.enterprise.api)
implementation(projects.libraries.architecture) implementation(projects.libraries.architecture)
implementation(projects.libraries.core) implementation(projects.libraries.core)
implementation(projects.libraries.di) implementation(projects.libraries.di)
@ -74,6 +75,7 @@ dependencies {
testImplementation(libs.test.turbine) testImplementation(libs.test.turbine)
testImplementation(libs.test.robolectric) testImplementation(libs.test.robolectric)
testImplementation(libs.kotlinx.collections.immutable) testImplementation(libs.kotlinx.collections.immutable)
testImplementation(projects.features.enterprise.test)
testImplementation(projects.libraries.matrix.test) testImplementation(projects.libraries.matrix.test)
testImplementation(projects.libraries.push.test) testImplementation(projects.libraries.push.test)
testImplementation(projects.libraries.pushstore.test) testImplementation(projects.libraries.pushstore.test)

View file

@ -0,0 +1,26 @@
/*
* Copyright 2025 New Vector Ltd.
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial
* Please see LICENSE files in the repository root for full details.
*/
package io.element.android.libraries.pushproviders.firebase
import com.squareup.anvil.annotations.ContributesBinding
import io.element.android.features.enterprise.api.EnterpriseService
import io.element.android.libraries.di.AppScope
import javax.inject.Inject
interface FirebaseGatewayProvider {
fun getFirebaseGateway(): String
}
@ContributesBinding(AppScope::class)
class DefaultFirebaseGatewayProvider @Inject constructor(
private val enterpriseService: EnterpriseService,
) : FirebaseGatewayProvider {
override fun getFirebaseGateway(): String {
return enterpriseService.firebasePushGateway() ?: FirebaseConfig.PUSHER_HTTP_URL
}
}

View file

@ -36,6 +36,7 @@ class DefaultFirebaseNewTokenHandler @Inject constructor(
private val userPushStoreFactory: UserPushStoreFactory, private val userPushStoreFactory: UserPushStoreFactory,
private val matrixClientProvider: MatrixClientProvider, private val matrixClientProvider: MatrixClientProvider,
private val firebaseStore: FirebaseStore, private val firebaseStore: FirebaseStore,
private val firebaseGatewayProvider: FirebaseGatewayProvider,
) : FirebaseNewTokenHandler { ) : FirebaseNewTokenHandler {
override suspend fun handle(firebaseToken: String) { override suspend fun handle(firebaseToken: String) {
firebaseStore.storeFcmToken(firebaseToken) firebaseStore.storeFcmToken(firebaseToken)
@ -55,7 +56,7 @@ class DefaultFirebaseNewTokenHandler @Inject constructor(
.registerPusher( .registerPusher(
matrixClient = client, matrixClient = client,
pushKey = firebaseToken, pushKey = firebaseToken,
gateway = FirebaseConfig.PUSHER_HTTP_URL, gateway = firebaseGatewayProvider.getFirebaseGateway(),
) )
.onFailure { .onFailure {
Timber.tag(loggerTag.value).e(it, "Failed to register pusher for session $sessionId") Timber.tag(loggerTag.value).e(it, "Failed to register pusher for session $sessionId")

View file

@ -27,6 +27,7 @@ class FirebasePushProvider @Inject constructor(
private val pusherSubscriber: PusherSubscriber, private val pusherSubscriber: PusherSubscriber,
private val isPlayServiceAvailable: IsPlayServiceAvailable, private val isPlayServiceAvailable: IsPlayServiceAvailable,
private val firebaseTokenRotator: FirebaseTokenRotator, private val firebaseTokenRotator: FirebaseTokenRotator,
private val firebaseGatewayProvider: FirebaseGatewayProvider,
) : PushProvider { ) : PushProvider {
override val index = FirebaseConfig.INDEX override val index = FirebaseConfig.INDEX
override val name = FirebaseConfig.NAME override val name = FirebaseConfig.NAME
@ -48,7 +49,7 @@ class FirebasePushProvider @Inject constructor(
return pusherSubscriber.registerPusher( return pusherSubscriber.registerPusher(
matrixClient = matrixClient, matrixClient = matrixClient,
pushKey = pushKey, pushKey = pushKey,
gateway = FirebaseConfig.PUSHER_HTTP_URL, gateway = firebaseGatewayProvider.getFirebaseGateway(),
) )
} }
@ -60,7 +61,7 @@ class FirebasePushProvider @Inject constructor(
Timber.tag(loggerTag.value).w("Unable to unregister pusher, Firebase token is not known.") Timber.tag(loggerTag.value).w("Unable to unregister pusher, Firebase token is not known.")
Result.success(Unit) Result.success(Unit)
} else { } else {
pusherSubscriber.unregisterPusher(matrixClient, pushKey, FirebaseConfig.PUSHER_HTTP_URL) pusherSubscriber.unregisterPusher(matrixClient, pushKey, firebaseGatewayProvider.getFirebaseGateway())
} }
} }
@ -72,7 +73,7 @@ class FirebasePushProvider @Inject constructor(
override suspend fun getCurrentUserPushConfig(): CurrentUserPushConfig? { override suspend fun getCurrentUserPushConfig(): CurrentUserPushConfig? {
return firebaseStore.getFcmToken()?.let { fcmToken -> return firebaseStore.getFcmToken()?.let { fcmToken ->
CurrentUserPushConfig( CurrentUserPushConfig(
url = FirebaseConfig.PUSHER_HTTP_URL, url = firebaseGatewayProvider.getFirebaseGateway(),
pushKey = fcmToken pushKey = fcmToken
) )
} }

View file

@ -79,8 +79,8 @@ class DefaultFirebaseNewTokenHandlerTest {
registerPusherResult.assertions() registerPusherResult.assertions()
.isCalledExactly(2) .isCalledExactly(2)
.withSequence( .withSequence(
listOf(value(aMatrixClient1), value("aToken"), value(FirebaseConfig.PUSHER_HTTP_URL)), listOf(value(aMatrixClient1), value("aToken"), value(A_FIREBASE_GATEWAY)),
listOf(value(aMatrixClient3), value("aToken"), value(FirebaseConfig.PUSHER_HTTP_URL)), listOf(value(aMatrixClient3), value("aToken"), value(A_FIREBASE_GATEWAY)),
) )
} }
@ -130,7 +130,7 @@ class DefaultFirebaseNewTokenHandlerTest {
registerPusherResult.assertions() registerPusherResult.assertions()
registerPusherResult.assertions() registerPusherResult.assertions()
.isCalledOnce() .isCalledOnce()
.with(value(aMatrixClient1), value("aToken"), value(FirebaseConfig.PUSHER_HTTP_URL)) .with(value(aMatrixClient1), value("aToken"), value(A_FIREBASE_GATEWAY))
} }
private fun createDefaultFirebaseNewTokenHandler( private fun createDefaultFirebaseNewTokenHandler(
@ -139,13 +139,15 @@ class DefaultFirebaseNewTokenHandlerTest {
userPushStoreFactory: UserPushStoreFactory = FakeUserPushStoreFactory(), userPushStoreFactory: UserPushStoreFactory = FakeUserPushStoreFactory(),
matrixClientProvider: MatrixClientProvider = FakeMatrixClientProvider(), matrixClientProvider: MatrixClientProvider = FakeMatrixClientProvider(),
firebaseStore: FirebaseStore = InMemoryFirebaseStore(), firebaseStore: FirebaseStore = InMemoryFirebaseStore(),
firebaseGatewayProvider: FirebaseGatewayProvider = FakeFirebaseGatewayProvider(),
): FirebaseNewTokenHandler { ): FirebaseNewTokenHandler {
return DefaultFirebaseNewTokenHandler( return DefaultFirebaseNewTokenHandler(
pusherSubscriber = pusherSubscriber, pusherSubscriber = pusherSubscriber,
sessionStore = sessionStore, sessionStore = sessionStore,
userPushStoreFactory = userPushStoreFactory, userPushStoreFactory = userPushStoreFactory,
matrixClientProvider = matrixClientProvider, matrixClientProvider = matrixClientProvider,
firebaseStore = firebaseStore firebaseStore = firebaseStore,
firebaseGatewayProvider = firebaseGatewayProvider,
) )
} }
} }

View file

@ -0,0 +1,16 @@
/*
* Copyright 2025 New Vector Ltd.
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial
* Please see LICENSE files in the repository root for full details.
*/
package io.element.android.libraries.pushproviders.firebase
const val A_FIREBASE_GATEWAY = "aGateway"
class FakeFirebaseGatewayProvider(
private val firebaseGatewayResult: () -> String = { A_FIREBASE_GATEWAY }
) : FirebaseGatewayProvider {
override fun getFirebaseGateway() = firebaseGatewayResult()
}

View file

@ -70,7 +70,7 @@ class FirebasePushProviderTest {
assertThat(result).isEqualTo(Result.success(Unit)) assertThat(result).isEqualTo(Result.success(Unit))
registerPusherResultLambda.assertions() registerPusherResultLambda.assertions()
.isCalledOnce() .isCalledOnce()
.with(value(matrixClient), value("aToken"), value(FirebaseConfig.PUSHER_HTTP_URL)) .with(value(matrixClient), value("aToken"), value(A_FIREBASE_GATEWAY))
} }
@Test @Test
@ -117,7 +117,7 @@ class FirebasePushProviderTest {
assertThat(result).isEqualTo(Result.success(Unit)) assertThat(result).isEqualTo(Result.success(Unit))
unregisterPusherResultLambda.assertions() unregisterPusherResultLambda.assertions()
.isCalledOnce() .isCalledOnce()
.with(value(matrixClient), value("aToken"), value(FirebaseConfig.PUSHER_HTTP_URL)) .with(value(matrixClient), value("aToken"), value(A_FIREBASE_GATEWAY))
} }
@Test @Test
@ -164,7 +164,7 @@ class FirebasePushProviderTest {
), ),
) )
val result = firebasePushProvider.getCurrentUserPushConfig() val result = firebasePushProvider.getCurrentUserPushConfig()
assertThat(result).isEqualTo(CurrentUserPushConfig(FirebaseConfig.PUSHER_HTTP_URL, "aToken")) assertThat(result).isEqualTo(CurrentUserPushConfig(A_FIREBASE_GATEWAY, "aToken"))
} }
@Test @Test
@ -194,12 +194,14 @@ class FirebasePushProviderTest {
pusherSubscriber: PusherSubscriber = FakePusherSubscriber(), pusherSubscriber: PusherSubscriber = FakePusherSubscriber(),
isPlayServiceAvailable: IsPlayServiceAvailable = FakeIsPlayServiceAvailable(false), isPlayServiceAvailable: IsPlayServiceAvailable = FakeIsPlayServiceAvailable(false),
firebaseTokenRotator: FirebaseTokenRotator = FakeFirebaseTokenRotator(), firebaseTokenRotator: FirebaseTokenRotator = FakeFirebaseTokenRotator(),
firebaseGatewayProvider: FirebaseGatewayProvider = FakeFirebaseGatewayProvider()
): FirebasePushProvider { ): FirebasePushProvider {
return FirebasePushProvider( return FirebasePushProvider(
firebaseStore = firebaseStore, firebaseStore = firebaseStore,
pusherSubscriber = pusherSubscriber, pusherSubscriber = pusherSubscriber,
isPlayServiceAvailable = isPlayServiceAvailable, isPlayServiceAvailable = isPlayServiceAvailable,
firebaseTokenRotator = firebaseTokenRotator, firebaseTokenRotator = firebaseTokenRotator,
firebaseGatewayProvider = firebaseGatewayProvider,
) )
} }
} }

View file

@ -19,6 +19,7 @@ setupAnvil()
dependencies { dependencies {
implementation(libs.dagger) implementation(libs.dagger)
implementation(projects.features.enterprise.api)
implementation(projects.libraries.androidutils) implementation(projects.libraries.androidutils)
implementation(projects.libraries.core) implementation(projects.libraries.core)
implementation(projects.libraries.matrix.api) implementation(projects.libraries.matrix.api)
@ -49,6 +50,7 @@ dependencies {
testImplementation(libs.test.truth) testImplementation(libs.test.truth)
testImplementation(libs.test.turbine) testImplementation(libs.test.turbine)
testImplementation(libs.kotlinx.collections.immutable) testImplementation(libs.kotlinx.collections.immutable)
testImplementation(projects.features.enterprise.test)
testImplementation(projects.libraries.matrix.test) testImplementation(projects.libraries.matrix.test)
testImplementation(projects.libraries.push.test) testImplementation(projects.libraries.push.test)
testImplementation(projects.libraries.pushproviders.test) testImplementation(projects.libraries.pushproviders.test)

View file

@ -0,0 +1,26 @@
/*
* Copyright 2025 New Vector Ltd.
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial
* Please see LICENSE files in the repository root for full details.
*/
package io.element.android.libraries.pushproviders.unifiedpush
import com.squareup.anvil.annotations.ContributesBinding
import io.element.android.features.enterprise.api.EnterpriseService
import io.element.android.libraries.di.AppScope
import javax.inject.Inject
interface DefaultPushGatewayHttpUrlProvider {
fun provide(): String
}
@ContributesBinding(AppScope::class)
class DefaultDefaultPushGatewayHttpUrlProvider @Inject constructor(
private val enterpriseService: EnterpriseService,
) : DefaultPushGatewayHttpUrlProvider {
override fun provide(): String {
return enterpriseService.unifiedPushDefaultPushGateway() ?: UnifiedPushConfig.DEFAULT_PUSH_GATEWAY_HTTP_URL
}
}

View file

@ -21,6 +21,7 @@ interface UnifiedPushGatewayUrlResolver {
@ContributesBinding(AppScope::class) @ContributesBinding(AppScope::class)
class DefaultUnifiedPushGatewayUrlResolver @Inject constructor( class DefaultUnifiedPushGatewayUrlResolver @Inject constructor(
private val unifiedPushStore: UnifiedPushStore, private val unifiedPushStore: UnifiedPushStore,
private val defaultPushGatewayHttpUrlProvider: DefaultPushGatewayHttpUrlProvider,
) : UnifiedPushGatewayUrlResolver { ) : UnifiedPushGatewayUrlResolver {
override fun resolve( override fun resolve(
gatewayResult: UnifiedPushGatewayResolverResult, gatewayResult: UnifiedPushGatewayResolverResult,
@ -33,7 +34,7 @@ class DefaultUnifiedPushGatewayUrlResolver @Inject constructor(
} }
UnifiedPushGatewayResolverResult.ErrorInvalidUrl, UnifiedPushGatewayResolverResult.ErrorInvalidUrl,
UnifiedPushGatewayResolverResult.NoMatrixGateway -> { UnifiedPushGatewayResolverResult.NoMatrixGateway -> {
UnifiedPushConfig.DEFAULT_PUSH_GATEWAY_HTTP_URL defaultPushGatewayHttpUrlProvider.provide()
} }
is UnifiedPushGatewayResolverResult.Success -> { is UnifiedPushGatewayResolverResult.Success -> {
gatewayResult.gateway gatewayResult.gateway

View file

@ -18,7 +18,7 @@ class DefaultUnifiedPushGatewayUrlResolverTest {
gatewayResult = UnifiedPushGatewayResolverResult.ErrorInvalidUrl, gatewayResult = UnifiedPushGatewayResolverResult.ErrorInvalidUrl,
instance = "", instance = "",
) )
assertThat(result).isEqualTo(UnifiedPushConfig.DEFAULT_PUSH_GATEWAY_HTTP_URL) assertThat(result).isEqualTo(A_UNIFIED_PUSH_GATEWAY)
} }
@Test @Test
@ -28,7 +28,7 @@ class DefaultUnifiedPushGatewayUrlResolverTest {
gatewayResult = UnifiedPushGatewayResolverResult.NoMatrixGateway, gatewayResult = UnifiedPushGatewayResolverResult.NoMatrixGateway,
instance = "", instance = "",
) )
assertThat(result).isEqualTo(UnifiedPushConfig.DEFAULT_PUSH_GATEWAY_HTTP_URL) assertThat(result).isEqualTo(A_UNIFIED_PUSH_GATEWAY)
} }
@Test @Test
@ -77,7 +77,9 @@ class DefaultUnifiedPushGatewayUrlResolverTest {
private fun createDefaultUnifiedPushGatewayUrlResolver( private fun createDefaultUnifiedPushGatewayUrlResolver(
unifiedPushStore: UnifiedPushStore = FakeUnifiedPushStore(), unifiedPushStore: UnifiedPushStore = FakeUnifiedPushStore(),
defaultPushGatewayHttpUrlProvider: DefaultPushGatewayHttpUrlProvider = FakeDefaultPushGatewayHttpUrlProvider(),
) = DefaultUnifiedPushGatewayUrlResolver( ) = DefaultUnifiedPushGatewayUrlResolver(
unifiedPushStore = unifiedPushStore, unifiedPushStore = unifiedPushStore,
defaultPushGatewayHttpUrlProvider = defaultPushGatewayHttpUrlProvider,
) )
} }

View file

@ -0,0 +1,18 @@
/*
* Copyright 2025 New Vector Ltd.
*
* SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-Element-Commercial
* Please see LICENSE files in the repository root for full details.
*/
package io.element.android.libraries.pushproviders.unifiedpush
const val A_UNIFIED_PUSH_GATEWAY = "aGateway"
class FakeDefaultPushGatewayHttpUrlProvider(
private val provideResult: () -> String = { A_UNIFIED_PUSH_GATEWAY }
) : DefaultPushGatewayHttpUrlProvider {
override fun provide(): String {
return provideResult()
}
}