diff --git a/.github/badges/branches.svg b/.github/badges/branches.svg index ebe2b57eb..b40b47c62 100644 --- a/.github/badges/branches.svg +++ b/.github/badges/branches.svg @@ -1 +1 @@ -branches38.8% \ No newline at end of file +branches39.9% \ No newline at end of file diff --git a/.github/badges/jacoco.svg b/.github/badges/jacoco.svg index 00b798798..dfeff7360 100644 --- a/.github/badges/jacoco.svg +++ b/.github/badges/jacoco.svg @@ -1 +1 @@ -coverage47.8% \ No newline at end of file +coverage49.6% \ No newline at end of file diff --git a/superwall/src/androidTest/java/com/superwall/sdk/config/ConfigManagerInstrumentedTest.kt b/superwall/src/androidTest/java/com/superwall/sdk/config/ConfigManagerInstrumentedTest.kt index ecedcb0f2..ff5667ff6 100644 --- a/superwall/src/androidTest/java/com/superwall/sdk/config/ConfigManagerInstrumentedTest.kt +++ b/superwall/src/androidTest/java/com/superwall/sdk/config/ConfigManagerInstrumentedTest.kt @@ -1,6 +1,5 @@ package com.superwall.sdk.config -import And import Given import Then import When @@ -10,7 +9,6 @@ import androidx.test.ext.junit.runners.AndroidJUnit4 import androidx.test.platform.app.InstrumentationRegistry import com.superwall.sdk.Superwall import com.superwall.sdk.analytics.Tier -import com.superwall.sdk.config.models.ConfigState import com.superwall.sdk.config.options.SuperwallOptions import com.superwall.sdk.dependencies.DependencyContainer import com.superwall.sdk.misc.Either diff --git a/superwall/src/androidTest/java/com/superwall/sdk/paywall/view/OffMainViewConstructionTest.kt b/superwall/src/androidTest/java/com/superwall/sdk/paywall/view/OffMainViewConstructionTest.kt new file mode 100644 index 000000000..657ffe9a4 --- /dev/null +++ b/superwall/src/androidTest/java/com/superwall/sdk/paywall/view/OffMainViewConstructionTest.kt @@ -0,0 +1,102 @@ +package com.superwall.sdk.paywall.view + +import android.graphics.Bitmap +import android.graphics.Canvas +import android.os.Looper +import android.view.View +import androidx.test.ext.junit.runners.AndroidJUnit4 +import androidx.test.platform.app.InstrumentationRegistry +import com.superwall.sdk.misc.ActivityProvider +import com.superwall.sdk.misc.primitives.SequentialActor +import com.superwall.sdk.network.device.DeviceHelper +import com.superwall.sdk.paywall.manager.PaywallCacheState +import com.superwall.sdk.paywall.manager.PaywallViewCache +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.SupervisorJob +import kotlinx.coroutines.cancel +import org.junit.Assert.assertNull +import org.junit.Assert.assertTrue +import org.junit.Test +import org.junit.runner.RunWith +import java.util.concurrent.ConcurrentHashMap + +/** + * PaywallViewCache builds the shared loading and shimmer views on its actor's + * IO thread, which has no Looper. These tests run on a real device to prove + * that construction is safe there and that the views still attach, draw and + * animate once handed to the main thread. + */ +@RunWith(AndroidJUnit4::class) +class OffMainViewConstructionTest { + private val instrumentation = InstrumentationRegistry.getInstrumentation() + private val ctx = instrumentation.targetContext + + private fun exerciseOnMain(vararg views: View) { + instrumentation.runOnMainSync { + views.forEach { view -> + view.measure( + View.MeasureSpec.makeMeasureSpec(400, View.MeasureSpec.EXACTLY), + View.MeasureSpec.makeMeasureSpec(800, View.MeasureSpec.EXACTLY), + ) + view.layout(0, 0, 400, 800) + (view as? PaywallShimmerView)?.showShimmer() + (view as? PaywallPurchaseLoadingView)?.showLoading() + view.draw(Canvas(Bitmap.createBitmap(400, 800, Bitmap.Config.ARGB_8888))) + (view as? PaywallShimmerView)?.hideShimmer() + } + } + } + + @Test + fun loadingAndShimmerCanBeBuiltOnAThreadWithoutALooper() { + var error: Throwable? = null + var hadLooper = true + var loading: LoadingView? = null + var shimmer: ShimmerView? = null + val thread = + Thread { + hadLooper = Looper.myLooper() != null + try { + loading = LoadingView(ctx, loadingColor = android.R.color.black) + shimmer = ShimmerView(ctx) + } catch (t: Throwable) { + error = t + } + } + thread.start() + thread.join() + + assertTrue("background thread must not have a Looper", !hadLooper) + assertNull("construction off main threw: $error", error) + exerciseOnMain(loading!!, shimmer!!) + } + + @Test + fun cacheAcquireFromMainReturnsUsableViewsWithoutDeadlocking() { + val scope = CoroutineScope(SupervisorJob() + Dispatchers.IO) + try { + val cache = + PaywallViewCache( + ctx, + object : ViewStorage { + override val views = ConcurrentHashMap() + }, + mockk { every { getCurrentActivity() } returns null }, + mockk { every { locale } returns "en_US" }, + actor = SequentialActor(PaywallCacheState(), scope), + ) + var loading: PaywallPurchaseLoadingView? = null + var shimmer: PaywallShimmerView? = null + instrumentation.runOnMainSync { + loading = cache.acquireLoadingView() + shimmer = cache.acquireShimmerView() + } + exerciseOnMain(loading as View, shimmer as View) + } finally { + scope.cancel() + } + } +} diff --git a/superwall/src/main/java/com/superwall/sdk/Superwall.kt b/superwall/src/main/java/com/superwall/sdk/Superwall.kt index cdcc333c4..bcec483df 100644 --- a/superwall/src/main/java/com/superwall/sdk/Superwall.kt +++ b/superwall/src/main/java/com/superwall/sdk/Superwall.kt @@ -15,7 +15,7 @@ import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent.* import com.superwall.sdk.analytics.superwall.SuperwallEventInfo import com.superwall.sdk.billing.toInternalResult -import com.superwall.sdk.config.models.ConfigState +import com.superwall.sdk.config.ConfigState import com.superwall.sdk.config.models.ConfigurationStatus import com.superwall.sdk.config.options.EventTrackingBehavior import com.superwall.sdk.config.options.SuperwallOptions @@ -98,7 +98,6 @@ import kotlinx.coroutines.flow.SharedFlow import kotlinx.coroutines.flow.StateFlow import kotlinx.coroutines.flow.asSharedFlow import kotlinx.coroutines.flow.distinctUntilChanged -import kotlinx.coroutines.flow.drop import kotlinx.coroutines.flow.filter import kotlinx.coroutines.flow.filterNotNull import kotlinx.coroutines.flow.map @@ -776,12 +775,15 @@ class Superwall( else -> old::class == new::class } } - .drop(1) // Drops the cached/initial emission - .collect { newValue -> + // Pair each status with the one before it. Entitlements persists the + // new status before this collector runs, so storage can't supply `from`. + .scan?>(null) { previous, newStatus -> + Pair(previous?.second, newStatus) + }.filterNotNull() + .filter { it.first != null } // Drops the cached/initial emission + .collect { (previous, newValue) -> // Save and handle the new value - val oldValue = - dependencyContainer.storage.read(StoredSubscriptionStatus) - ?: SubscriptionStatus.Unknown + val oldValue = previous ?: SubscriptionStatus.Unknown dependencyContainer.storage.write(StoredSubscriptionStatus, newValue) dependencyContainer.delegateAdapter.subscriptionStatusDidChange( oldValue, diff --git a/superwall/src/main/java/com/superwall/sdk/analytics/internal/Tracking.kt b/superwall/src/main/java/com/superwall/sdk/analytics/internal/Tracking.kt index 706b81c49..e896b1725 100644 --- a/superwall/src/main/java/com/superwall/sdk/analytics/internal/Tracking.kt +++ b/superwall/src/main/java/com/superwall/sdk/analytics/internal/Tracking.kt @@ -4,7 +4,7 @@ import com.superwall.sdk.Superwall import com.superwall.sdk.analytics.internal.trackable.Trackable import com.superwall.sdk.analytics.internal.trackable.TrackableSuperwallEvent import com.superwall.sdk.analytics.superwall.SuperwallEventInfo -import com.superwall.sdk.config.models.ConfigState +import com.superwall.sdk.config.ConfigState import com.superwall.sdk.logger.LogLevel import com.superwall.sdk.logger.LogScope import com.superwall.sdk.logger.Logger diff --git a/superwall/src/main/java/com/superwall/sdk/analytics/session/AppSessionManager.kt b/superwall/src/main/java/com/superwall/sdk/analytics/session/AppSessionManager.kt index 417786d0c..feb063433 100644 --- a/superwall/src/main/java/com/superwall/sdk/analytics/session/AppSessionManager.kt +++ b/superwall/src/main/java/com/superwall/sdk/analytics/session/AppSessionManager.kt @@ -6,7 +6,7 @@ import com.superwall.sdk.Superwall import com.superwall.sdk.analytics.internal.track import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent import com.superwall.sdk.config.ConfigManager -import com.superwall.sdk.config.models.getConfig +import com.superwall.sdk.config.getConfig import com.superwall.sdk.dependencies.DeviceHelperFactory import com.superwall.sdk.dependencies.UserAttributesEventFactory import com.superwall.sdk.misc.IOScope diff --git a/superwall/src/main/java/com/superwall/sdk/config/ConfigContext.kt b/superwall/src/main/java/com/superwall/sdk/config/ConfigContext.kt index 368d8a100..7327c766c 100644 --- a/superwall/src/main/java/com/superwall/sdk/config/ConfigContext.kt +++ b/superwall/src/main/java/com/superwall/sdk/config/ConfigContext.kt @@ -1,12 +1,9 @@ package com.superwall.sdk.config import android.content.Context -import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent -import com.superwall.sdk.config.models.ConfigState import com.superwall.sdk.config.options.SuperwallOptions import com.superwall.sdk.identity.IdentityManager import com.superwall.sdk.misc.primitives.BaseContext -import com.superwall.sdk.models.config.Config import com.superwall.sdk.models.entitlements.SubscriptionStatus import com.superwall.sdk.models.triggers.Trigger import com.superwall.sdk.network.SuperwallAPI @@ -33,7 +30,6 @@ interface ConfigContext : BaseContext { val identityManager: (() -> IdentityManager)? val setSubscriptionStatus: ((SubscriptionStatus) -> Unit)? val awaitUtilNetwork: suspend () -> Unit - val activateTestMode: suspend (config: Config, justActivated: Boolean) -> Unit fun setTriggers(triggers: Map) } diff --git a/superwall/src/main/java/com/superwall/sdk/config/ConfigManager.kt b/superwall/src/main/java/com/superwall/sdk/config/ConfigManager.kt index bc377ad6f..e571d58fe 100644 --- a/superwall/src/main/java/com/superwall/sdk/config/ConfigManager.kt +++ b/superwall/src/main/java/com/superwall/sdk/config/ConfigManager.kt @@ -1,10 +1,7 @@ package com.superwall.sdk.config import android.content.Context -import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent import com.superwall.sdk.analytics.internal.trackable.TrackableSuperwallEvent -import com.superwall.sdk.config.models.ConfigState -import com.superwall.sdk.config.models.getConfig import com.superwall.sdk.config.options.SuperwallOptions import com.superwall.sdk.dependencies.DeviceHelperFactory import com.superwall.sdk.dependencies.DeviceInfoFactory @@ -36,7 +33,6 @@ import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.StateFlow import kotlinx.coroutines.flow.mapNotNull import kotlinx.coroutines.flow.take -import kotlinx.coroutines.launch open class ConfigManager( override val context: Context, @@ -59,7 +55,6 @@ open class ConfigManager( override val awaitUtilNetwork: suspend () -> Unit = { context.awaitUntilNetworkExists() }, - override val activateTestMode: suspend (Config, Boolean) -> Unit = { _, _ -> }, override val actor: StateActor, ) : ConfigContext { interface Factory : diff --git a/superwall/src/main/java/com/superwall/sdk/config/models/ConfigState.kt b/superwall/src/main/java/com/superwall/sdk/config/ConfigState.kt similarity index 98% rename from superwall/src/main/java/com/superwall/sdk/config/models/ConfigState.kt rename to superwall/src/main/java/com/superwall/sdk/config/ConfigState.kt index 0d5b7ad21..989720472 100644 --- a/superwall/src/main/java/com/superwall/sdk/config/models/ConfigState.kt +++ b/superwall/src/main/java/com/superwall/sdk/config/ConfigState.kt @@ -1,15 +1,11 @@ -package com.superwall.sdk.config.models +package com.superwall.sdk.config import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent -import com.superwall.sdk.config.ConfigContext -import com.superwall.sdk.config.ConfigLogic -import com.superwall.sdk.config.PaywallPreload import com.superwall.sdk.config.options.computedShouldPreload import com.superwall.sdk.logger.LogLevel import com.superwall.sdk.logger.LogScope import com.superwall.sdk.logger.Logger import com.superwall.sdk.misc.Either -import com.superwall.sdk.misc.awaitFirstValidConfig import com.superwall.sdk.misc.fold import com.superwall.sdk.misc.into import com.superwall.sdk.misc.onError @@ -326,7 +322,7 @@ sealed class ConfigState { manager.setOverriddenSubscriptionStatus(defaultStatus) entitlements.setSubscriptionStatus(defaultStatus) } - scope.launch { activateTestMode(config, testModeJustActivated) } + scope.launch { manager.activate(config, testModeJustActivated) } } else { if (wasTestMode) { manager?.clearTestModeState() @@ -357,7 +353,7 @@ sealed class ConfigState { manager.clearTestModeState() setSubscriptionStatus?.invoke(SubscriptionStatus.Inactive) } else if (!wasTestMode && isNowTestMode) { - scope.launch { activateTestMode(config, true) } + scope.launch { manager.activate(config, justActivated = true) } } }) diff --git a/superwall/src/main/java/com/superwall/sdk/debug/DebugView.kt b/superwall/src/main/java/com/superwall/sdk/debug/DebugView.kt index 618663d74..2161ab9ac 100644 --- a/superwall/src/main/java/com/superwall/sdk/debug/DebugView.kt +++ b/superwall/src/main/java/com/superwall/sdk/debug/DebugView.kt @@ -924,7 +924,7 @@ internal class DebugViewActivity : AppCompatActivity() { ) { val key = UUID.randomUUID().toString() Superwall.instance.dependencyContainer - .makeViewStore() + .makeViewRegistry() .storeView(key, view) val intent = @@ -962,7 +962,7 @@ internal class DebugViewActivity : AppCompatActivity() { } val view = Superwall.instance.dependencyContainer - .makeViewStore() + .makeViewRegistry() .retrieveView(key) ?: run { finish() // Close the activity if the view associated with the key is not found return diff --git a/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt b/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt index da7d94c79..cd4eafc58 100644 --- a/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt +++ b/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt @@ -28,6 +28,7 @@ import com.superwall.sdk.billing.GoogleBillingWrapper import com.superwall.sdk.config.Assignments import com.superwall.sdk.config.ConfigLogic import com.superwall.sdk.config.ConfigManager +import com.superwall.sdk.config.ConfigState import com.superwall.sdk.config.PaywallPreload import com.superwall.sdk.config.options.SuperwallOptions import com.superwall.sdk.customer.CustomerInfoManager @@ -52,11 +53,11 @@ import com.superwall.sdk.misc.AppLifecycleObserver import com.superwall.sdk.misc.CurrentActivityTracker import com.superwall.sdk.misc.IOScope import com.superwall.sdk.misc.MainScope -import com.superwall.sdk.misc.primitives.DebugInterceptor import com.superwall.sdk.misc.primitives.SequentialActor import com.superwall.sdk.misc.sha256Hex import com.superwall.sdk.models.config.ComputedPropertyRequest import com.superwall.sdk.models.config.FeatureFlags +import com.superwall.sdk.models.entitlements.Entitlement import com.superwall.sdk.models.entitlements.SubscriptionStatus import com.superwall.sdk.models.entitlements.TransactionReceipt import com.superwall.sdk.models.events.EventData @@ -75,8 +76,10 @@ import com.superwall.sdk.network.SubscriptionService import com.superwall.sdk.network.device.DeviceHelper import com.superwall.sdk.network.device.DeviceInfo import com.superwall.sdk.network.session.CustomHttpUrlConnection +import com.superwall.sdk.paywall.manager.PaywallCacheState import com.superwall.sdk.paywall.manager.PaywallManager import com.superwall.sdk.paywall.manager.PaywallViewCache +import com.superwall.sdk.paywall.manager.PaywallViewRegistry import com.superwall.sdk.paywall.presentation.CustomCallbackRegistry import com.superwall.sdk.paywall.presentation.PaywallInfo import com.superwall.sdk.paywall.presentation.dismiss @@ -278,7 +281,7 @@ class DependencyContainer( json = json(), _apiKey = apiKey ) - entitlements = Entitlements(storage) + entitlements = Entitlements(storage, actorScope = ioScope) val options = options ?: SuperwallOptions() testMode = TestMode( @@ -295,7 +298,8 @@ class DependencyContainer( else -> "https://superwall.com" } }, - track = { Superwall.instance.track(it) }, + tracker = { Superwall.instance.track(it) }, + ioScope = ioScope, ) testModeTransactionHandler = TestModeTransactionHandler( @@ -463,8 +467,8 @@ class DependencyContainer( // actions (fetch, refresh, reset, reevaluate test mode) through a single // FIFO queue, so applying a new config can never race with a variant pick. val configActor = - SequentialActor( - com.superwall.sdk.config.models.ConfigState.None, + SequentialActor( + ConfigState.None, ioScope, ) // DebugInterceptor.install(configActor, name = "Config") @@ -492,9 +496,6 @@ class DependencyContainer( setSubscriptionStatus = { status -> entitlements.setSubscriptionStatus(status) }, - activateTestMode = { config, justActivated -> - testMode.activate(config, justActivated) - }, actor = configActor, ) @@ -909,6 +910,7 @@ class DependencyContainer( activityProvider!!, deviceHelper, configManager.options.paywalls.loadingColor, + actor = SequentialActor(PaywallCacheState(), ioScope), ) override fun activePaywallId(): String? = @@ -1123,6 +1125,12 @@ class DependencyContainer( override fun makeViewStore(): ViewStorageViewModel = vmProvider[ViewStorageViewModel::class.java] + /** + * The only way Activities and debug UI should reach paywall views, so the + * cache and ViewStorage stay in sync. Internal because the registry is. + */ + internal fun makeViewRegistry(): PaywallViewRegistry = paywallManager.viewRegistry + private var _mainScope: MainScope? = null private var _ioScope: IOScope? = null @@ -1262,6 +1270,10 @@ class DependencyContainer( Superwall.instance.track(event) } + override fun setWebEntitlements(entitlements: Set) { + this.entitlements.setWebEntitlements(entitlements) + } + override fun internallySetSubscriptionStatus(status: SubscriptionStatus) { Superwall.instance.internallySetSubscriptionStatus(status) } diff --git a/superwall/src/main/java/com/superwall/sdk/misc/Config+AwaitFirstValidConfig.kt b/superwall/src/main/java/com/superwall/sdk/misc/Config+AwaitFirstValidConfig.kt index 539c46982..4fc130ff2 100644 --- a/superwall/src/main/java/com/superwall/sdk/misc/Config+AwaitFirstValidConfig.kt +++ b/superwall/src/main/java/com/superwall/sdk/misc/Config+AwaitFirstValidConfig.kt @@ -1,6 +1,6 @@ package com.superwall.sdk.misc -import com.superwall.sdk.config.models.ConfigState +import com.superwall.sdk.config.ConfigState import com.superwall.sdk.models.config.Config import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.filterIsInstance diff --git a/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallManager.kt b/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallManager.kt index a4304f7c0..cb60b289f 100644 --- a/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallManager.kt +++ b/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallManager.kt @@ -38,13 +38,16 @@ class PaywallManager( return _cache!! } + /** Narrow view of the cache for Activities and debug UI. */ + internal val viewRegistry: PaywallViewRegistry by lazy { cache.asRegistry() } + private fun createCache(): PaywallViewCache { val cache: PaywallViewCache = factory.makeCache() _cache = cache return cache } - fun removePaywallView(identifier: PaywallIdentifier) { + suspend fun removePaywallView(identifier: PaywallIdentifier) { cache.removePaywallView(identifier) } diff --git a/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallViewCache.kt b/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallViewCache.kt index 1319ab40c..2a11250f9 100644 --- a/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallViewCache.kt +++ b/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallViewCache.kt @@ -1,8 +1,13 @@ package com.superwall.sdk.paywall.manager import android.content.Context +import android.view.View import androidx.annotation.ColorRes import com.superwall.sdk.misc.ActivityProvider +import com.superwall.sdk.misc.primitives.Reducer +import com.superwall.sdk.misc.primitives.SequentialActor +import com.superwall.sdk.misc.primitives.StoreContext +import com.superwall.sdk.misc.primitives.TypedAction import com.superwall.sdk.models.paywall.PaywallIdentifier import com.superwall.sdk.network.device.DeviceHelper import com.superwall.sdk.paywall.view.LoadingView @@ -11,92 +16,244 @@ import com.superwall.sdk.paywall.view.PaywallShimmerView import com.superwall.sdk.paywall.view.PaywallView import com.superwall.sdk.paywall.view.ShimmerView import com.superwall.sdk.paywall.view.ViewStorage +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.runBlocking -class PaywallViewCache( - private val appCtx: Context, - private val store: ViewStorage, - private val activityProvider: ActivityProvider, - private val deviceHelper: DeviceHelper, - @ColorRes private val loadingColor: Int? = null, +/** + * Source-of-truth state for the paywall view cache. + * + * Mirrors the underlying [ViewStorage] but is owned by a [SequentialActor] so + * mutations are FIFO-serialized and reads via `actor.state.value` always + * return a consistent snapshot. + */ +data class PaywallCacheState( + val views: Map = emptyMap(), + val activePaywallVcKey: String? = null, ) { - private val ctx: Context + val paywallViews: List + get() = views.values.filterIsInstance() + + fun viewAt(key: String): View? = views[key] + + val activePaywallView: PaywallView? + get() = activePaywallVcKey?.let { views[it] as? PaywallView } + + internal sealed class Updates( + override val reduce: (PaywallCacheState) -> PaywallCacheState, + ) : Reducer { + data class StoreView( + val key: String, + val view: View, + ) : Updates({ it.copy(views = it.views + (key to view)) }) + + data class RemoveView( + val key: String, + ) : Updates({ it.copy(views = it.views - key) }) + + data class SetActiveKey( + val key: String?, + ) : Updates({ it.copy(activePaywallVcKey = key) }) + + data class RemoveViews( + val keys: Set, + ) : Updates({ it.copy(views = it.views - keys) }) + + data class Hydrate( + val views: Map, + ) : Updates({ it.copy(views = views) }) + } + + internal sealed class Actions( + override val execute: suspend PaywallCacheContext.() -> Unit, + ) : TypedAction { + /** Atomically write a paywall view to state and viewStorage. */ + data class Save( + val identifier: PaywallIdentifier, + val view: PaywallView, + ) : Actions({ + val key = PaywallCacheLogic.key(identifier, deviceHelper.locale) + viewStorage.storeView(key, view) + update(Updates.StoreView(key, view)) + }) + + data class Remove( + val identifier: PaywallIdentifier, + ) : Actions({ + val key = PaywallCacheLogic.key(identifier, deviceHelper.locale) + viewStorage.removeView(key) + update(Updates.RemoveView(key)) + }) + + /** + * Evicts every key except the active paywall. Removes exactly the keys + * it evicted from viewStorage, so a view stored concurrently through + * [PaywallViewCache.storeView] survives in both places. + */ + object RemoveAllExceptActive : Actions({ + val active = state.value.activePaywallVcKey + val evicted = state.value.views.keys.filterTo(mutableSetOf()) { it != active } + evicted.forEach { viewStorage.removeView(it) } + update(Updates.RemoveViews(evicted)) + }) + + /** + * Get-or-create the LoadingView. Atomic: only one factory invocation + * across concurrent callers because actions are FIFO-serialized. + * + * The factory runs on the actor's consumer thread and must not dispatch + * to `Dispatchers.Main`: callers block on the result via `runBlocking`, + * usually from the main thread, so a main hop here would deadlock. + */ + data class EnsureLoadingView( + val factory: () -> PaywallPurchaseLoadingView, + ) : Actions({ + if (state.value.views[LoadingView.TAG] !is PaywallPurchaseLoadingView) { + val v = factory() + viewStorage.storeView(LoadingView.TAG, v as View) + update(Updates.StoreView(LoadingView.TAG, v)) + } + }) + + data class EnsureShimmerView( + val factory: () -> PaywallShimmerView, + ) : Actions({ + if (state.value.views[ShimmerView.TAG] !is PaywallShimmerView) { + val v = factory() + viewStorage.storeView(ShimmerView.TAG, v as View) + update(Updates.StoreView(ShimmerView.TAG, v)) + } + }) + } +} + +/** + * Dependencies available to [PaywallCacheState.Actions]. + * + * [PaywallViewCache] implements this directly — actions receive `this` as + * their context, with no intermediate object. + */ +interface PaywallCacheContext : StoreContext { + val viewStorage: ViewStorage + val deviceHelper: DeviceHelper + val activityProvider: ActivityProvider + val appCtx: Context + + val ctx: Context get() = activityProvider.getCurrentActivity() ?: appCtx +} - @Volatile - private var _activePaywallVcKey: String? = null - private val loadingView: LoadingView = LoadingView(context = ctx, loadingColor = loadingColor) - private val shimmerView: ShimmerView = ShimmerView(context = ctx) +/** + * Cache for paywall, loading, and shimmer views. + * + * State is owned by a [SequentialActor] — every mutation is enqueued through a + * single FIFO consumer, so `state.value` always reflects the latest committed + * data and there are no races between save/get, remove/save, or concurrent + * acquire calls. [ViewStorage] is kept as a write-through mirror because + * external readers (SuperwallPaywallActivity, DebugView) read it directly. + * Every write must go through this cache ([save], [storeView], [removeView], + * ...) so the two never disagree about which keys exist. + */ +class PaywallViewCache( + override val appCtx: Context, + override val viewStorage: ViewStorage, + override val activityProvider: ActivityProvider, + override val deviceHelper: DeviceHelper, + @ColorRes private val loadingColor: Int? = null, + override val actor: SequentialActor, +) : PaywallCacheContext { + override val scope: CoroutineScope get() = actor.scope init { - store.storeView(LoadingView.TAG, loadingView) - store.storeView(ShimmerView.TAG, shimmerView) + // Hydrate from any pre-existing entries in viewStorage (e.g. survived + // an Activity recreation via the ViewStorageViewModel). + val existing = viewStorage.views.toMap() + if (existing.isNotEmpty()) { + actor.update(PaywallCacheState.Updates.Hydrate(existing)) + } } - fun getAllPaywallViews(): List = store.all().filterIsInstance().toList() + val entries: Map + get() = state.value.views var activePaywallVcKey: String? - get() = _activePaywallVcKey + get() = state.value.activePaywallVcKey set(value) { - _activePaywallVcKey = value + actor.update(PaywallCacheState.Updates.SetActiveKey(value)) } val activePaywallView: PaywallView? - get() = _activePaywallVcKey?.let { store.retrieveView(it) as PaywallView? } + get() = state.value.activePaywallView - fun save( + fun getAllPaywallViews(): List = state.value.paywallViews + + fun getPaywallView(key: String): PaywallView? = state.value.viewAt(key) as? PaywallView + + suspend fun save( paywallView: PaywallView, identifier: PaywallIdentifier, ) { - store.storeView( - PaywallCacheLogic.key( - identifier, - locale = deviceHelper.locale, - ), - paywallView, - ) + immediate(PaywallCacheState.Actions.Save(identifier, paywallView)) } - fun acquireLoadingView(): PaywallPurchaseLoadingView { - return store.retrieveView(LoadingView.TAG)?.let { - it as PaywallPurchaseLoadingView - } ?: run { - val view = LoadingView(ctx, loadingColor = loadingColor) - store.storeView(LoadingView.TAG, view) - return view - } + suspend fun removePaywallView(identifier: PaywallIdentifier) { + immediate(PaywallCacheState.Actions.Remove(identifier)) } - fun acquireShimmerView(): PaywallShimmerView { - return store.retrieveView(ShimmerView.TAG)?.let { - it as PaywallShimmerView - } ?: run { - val view = ShimmerView(ctx) - store.storeView(ShimmerView.TAG, view) - return view - } + /** + * Stores a view under an arbitrary key (activity launch and restore keys, + * debug views). Synchronous because callers hand [key] to an Activity that + * reads [ViewStorage] as soon as it is created. Both writes are atomic map + * operations, so no queued action is needed for consistency. + */ + fun storeView( + key: String, + view: View, + ) { + viewStorage.storeView(key, view) + actor.update(PaywallCacheState.Updates.StoreView(key, view)) } - fun getPaywallView(key: String): PaywallView? = - try { - store.retrieveView(key) as PaywallView? - } catch (e: Throwable) { - null - } + /** Synchronous counterpart of [storeView]. */ + fun removeView(key: String) { + viewStorage.removeView(key) + actor.update(PaywallCacheState.Updates.RemoveView(key)) + } - fun removePaywallView(identifier: PaywallIdentifier) { - store.removeView( - PaywallCacheLogic.key( - identifier, - locale = deviceHelper.locale, - ), - ) + suspend fun removeAll() { + immediate(PaywallCacheState.Actions.RemoveAllExceptActive) } - fun removeAll() { - store.views.keys.forEach { key -> - if (key != _activePaywallVcKey) { - store.removeView(key) - } + /** + * Synchronous because [PaywallView.present] is non-suspend. + * + * Fast path: if state already holds the canonical view, return it without + * touching the actor queue. Slow path (cold start, or after [removeAll] + * evicted the tag): block on the actor's `immediate` so exactly one + * factory invocation happens across concurrent callers. The View is + * constructed on the actor thread; it is only attached to a hierarchy + * later, on the main thread, by [PaywallView]. + */ + fun acquireLoadingView(): PaywallPurchaseLoadingView { + (state.value.views[LoadingView.TAG] as? PaywallPurchaseLoadingView)?.let { return it } + return runBlocking { + immediate( + PaywallCacheState.Actions.EnsureLoadingView { + LoadingView(ctx, loadingColor = loadingColor) + }, + ) + state.value.views[LoadingView.TAG] as PaywallPurchaseLoadingView + } + } + + fun acquireShimmerView(): PaywallShimmerView { + (state.value.views[ShimmerView.TAG] as? PaywallShimmerView)?.let { return it } + return runBlocking { + immediate( + PaywallCacheState.Actions.EnsureShimmerView { + ShimmerView(ctx) + }, + ) + state.value.views[ShimmerView.TAG] as PaywallShimmerView } } } diff --git a/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallViewRegistry.kt b/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallViewRegistry.kt new file mode 100644 index 000000000..462dadc80 --- /dev/null +++ b/superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallViewRegistry.kt @@ -0,0 +1,46 @@ +package com.superwall.sdk.paywall.manager + +import android.view.View +import com.superwall.sdk.paywall.view.PaywallPurchaseLoadingView +import com.superwall.sdk.paywall.view.PaywallShimmerView + +/** + * What Activities and debug UI need from the paywall view cache: hand a view + * to an Activity by key, look it up again, and borrow the shared loading and + * shimmer views. + * + * Callers depend on this instead of [PaywallViewCache] or + * [com.superwall.sdk.paywall.view.ViewStorage] directly, so every write keeps + * the cache and the storage in sync and the cache can change underneath. + */ +internal interface PaywallViewRegistry { + fun storeView( + key: String, + view: View, + ) + + fun removeView(key: String) + + fun retrieveView(key: String): View? + + fun acquireLoadingView(): PaywallPurchaseLoadingView + + fun acquireShimmerView(): PaywallShimmerView +} + +internal fun PaywallViewCache.asRegistry(): PaywallViewRegistry = + object : PaywallViewRegistry { + override fun storeView( + key: String, + view: View, + ) = this@asRegistry.storeView(key, view) + + override fun removeView(key: String) = this@asRegistry.removeView(key) + + // Read ViewStorage: it is the copy that survives Activity recreation. + override fun retrieveView(key: String): View? = viewStorage.retrieveView(key) + + override fun acquireLoadingView() = this@asRegistry.acquireLoadingView() + + override fun acquireShimmerView() = this@asRegistry.acquireShimmerView() + } diff --git a/superwall/src/main/java/com/superwall/sdk/paywall/presentation/internal/operators/WaitForSubsStatusAndConfig.kt b/superwall/src/main/java/com/superwall/sdk/paywall/presentation/internal/operators/WaitForSubsStatusAndConfig.kt index 2f6c8e06a..5c122ce35 100644 --- a/superwall/src/main/java/com/superwall/sdk/paywall/presentation/internal/operators/WaitForSubsStatusAndConfig.kt +++ b/superwall/src/main/java/com/superwall/sdk/paywall/presentation/internal/operators/WaitForSubsStatusAndConfig.kt @@ -1,9 +1,8 @@ package com.superwall.sdk.paywall.presentation.internal.operators import com.superwall.sdk.Superwall -import com.superwall.sdk.analytics.internal.track import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent -import com.superwall.sdk.config.models.ConfigState +import com.superwall.sdk.config.ConfigState import com.superwall.sdk.dependencies.DependencyContainer import com.superwall.sdk.logger.LogLevel import com.superwall.sdk.logger.LogScope diff --git a/superwall/src/main/java/com/superwall/sdk/paywall/view/SuperwallPaywallActivity.kt b/superwall/src/main/java/com/superwall/sdk/paywall/view/SuperwallPaywallActivity.kt index c8332e5d9..246c8aa9b 100644 --- a/superwall/src/main/java/com/superwall/sdk/paywall/view/SuperwallPaywallActivity.kt +++ b/superwall/src/main/java/com/superwall/sdk/paywall/view/SuperwallPaywallActivity.kt @@ -158,7 +158,7 @@ class SuperwallPaywallActivity : AppCompatActivity() { return launchPaywallActivity(context, intent).onFailure { Superwall.instance.dependencyContainer - .makeViewStore() + .makeViewRegistry() .removeView(key) view.clearActivityLaunchState() } @@ -167,12 +167,13 @@ class SuperwallPaywallActivity : AppCompatActivity() { private fun PaywallView.prepareViewForDisplay(key: String) { webView.enableBackgroundRendering() webView.attach(this) - val viewStorageViewModel = Superwall.instance.dependencyContainer.makeViewStore() + val registry = Superwall.instance.dependencyContainer.makeViewRegistry() // If we started it directly and the view does not have shimmer and loading attached - // We set them up for this PaywallView + // We set them up for this PaywallView. Acquire through the cache rather than reading + // ViewStorage: the canonical views are created lazily, so they may not exist yet + // (getPaywall() + startWithView() without a prior present(), or after resetCache()). if (children.none { it is LoadingView || it is ShimmerView }) { - val loading = - (viewStorageViewModel.retrieveView(LoadingView.TAG) as LoadingView) + val loading = registry.acquireLoadingView() val style = state.paywall.presentation.style val shimmer = if (style is PaywallPresentationStyle.Popup) { @@ -184,12 +185,12 @@ class SuperwallPaywallActivity : AppCompatActivity() { ) } } else { - (viewStorageViewModel.retrieveView(ShimmerView.TAG) as ShimmerView) + registry.acquireShimmerView() } setupWith(shimmer, loading) } - viewStorageViewModel.storeView(key, this) + registry.storeView(key, this) } } @@ -241,9 +242,9 @@ class SuperwallPaywallActivity : AppCompatActivity() { return } - val viewStorageViewModel = + val viewRegistry = try { - Superwall.instance.dependencyContainer.makeViewStore() + Superwall.instance.dependencyContainer.makeViewRegistry() } catch (e: Exception) { Logger.debug( LogLevel.error, @@ -254,7 +255,7 @@ class SuperwallPaywallActivity : AppCompatActivity() { } val view = - viewStorageViewModel.retrieveView(key) as? PaywallView ?: run { + viewRegistry.retrieveView(key) as? PaywallView ?: run { Logger.debug( LogLevel.error, LogScope.paywallView, @@ -286,7 +287,7 @@ class SuperwallPaywallActivity : AppCompatActivity() { } // Store the view again with the same key for this activity - viewStorageViewModel.storeView(key, currentPaywallView) + viewRegistry.storeView(key, currentPaywallView) // Continue with normal activity setup using the restored view setupActivityWithView(currentPaywallView, presentationStyle) return @@ -844,7 +845,7 @@ class SuperwallPaywallActivity : AppCompatActivity() { if (pv != null) { ( Superwall.instance.dependencyContainer - .makeViewStore() + .makeViewRegistry() .retrieveView(pv) as? PaywallView? )?.cleanup() } diff --git a/superwall/src/main/java/com/superwall/sdk/store/Entitlements.kt b/superwall/src/main/java/com/superwall/sdk/store/Entitlements.kt index 294127c67..5555e10b0 100644 --- a/superwall/src/main/java/com/superwall/sdk/store/Entitlements.kt +++ b/superwall/src/main/java/com/superwall/sdk/store/Entitlements.kt @@ -1,11 +1,9 @@ package com.superwall.sdk.store -import com.superwall.sdk.billing.DecomposedProductIds -import com.superwall.sdk.models.customer.mergeEntitlementsPrioritized -import com.superwall.sdk.models.customer.toSet +import com.superwall.sdk.analytics.internal.trackable.TrackableSuperwallEvent +import com.superwall.sdk.misc.primitives.StateActor import com.superwall.sdk.models.entitlements.Entitlement import com.superwall.sdk.models.entitlements.SubscriptionStatus -import com.superwall.sdk.storage.LatestRedemptionResponse import com.superwall.sdk.storage.Storage import com.superwall.sdk.storage.StoredEntitlementsByProductId import com.superwall.sdk.storage.StoredSubscriptionStatus @@ -15,37 +13,32 @@ import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.flow.StateFlow import kotlinx.coroutines.flow.asStateFlow import kotlinx.coroutines.launch -import java.util.concurrent.ConcurrentHashMap /** - * A class that handles the Set of Entitlement objects retrieved from - * the Superwall dashboard. + * Facade over the entitlements state held in a [StateActor]. + * + * Implements [EntitlementsContext] directly — actions receive `this` as + * their context, eliminating the intermediate object. + * + * State mutations use [StateActor.update] (synchronous CAS, routed through + * interceptors) and are persisted immediately through [persist]. The initial + * state is rebuilt from storage by [createInitialEntitlementsState] before + * the actor starts, so cached status, product entitlements and web + * entitlements are available synchronously after construction. */ class Entitlements( - private val storage: Storage, - private val scope: CoroutineScope = CoroutineScope(Dispatchers.Default), -) { - val web: Set - get() = - storage - .read(LatestRedemptionResponse) - ?.customerInfo - ?.entitlements - ?.filter { it.isActive } - ?.toSet() ?: emptySet() - - // MARK: - Private Properties - internal val entitlementsByProduct = ConcurrentHashMap>() + override val storage: Storage, + actorScope: CoroutineScope = CoroutineScope(Dispatchers.Default), + override val actor: StateActor = + StateActor(createInitialEntitlementsState(storage), actorScope), + override val tracker: suspend (TrackableSuperwallEvent) -> Unit = {}, +) : EntitlementsContext { + override val scope: CoroutineScope = actorScope - /** - * Returns a snapshot of all entitlements by product ID. - * Used when loading purchases to enrich entitlements with transaction data. - */ - val entitlementsByProductId: Map> - get() = entitlementsByProduct.toMap() + // -- Status flow (kept in sync with actor state for external collection) -- private val _status: MutableStateFlow = - MutableStateFlow(SubscriptionStatus.Unknown) + MutableStateFlow(actor.state.value.status) /** * A StateFlow of the entitlement status of the user. Set this using @@ -56,30 +49,40 @@ class Entitlements( val status: StateFlow get() = _status.asStateFlow() - // MARK: - Backing Fields + init { + scope.launch { + actor.state.collect { _status.value = it.status } + } + } + + private val snapshot get() = actor.state.value /** - * Internal backing variable that is set only via setSubscriptionStatus + * Active web entitlements from the latest redemption response. + * Updated by [WebPaywallRedeemer] through [setWebEntitlements]. */ - private var backingActive: MutableSet = mutableSetOf() + val web: Set + get() = snapshot.webEntitlements - private val _all = mutableSetOf() - private val _activeDeviceEntitlements = mutableSetOf() - private val _inactive = _all.subtract(backingActive).toMutableSet() - // MARK: - Public Properties + /** + * Returns a snapshot of all entitlements by product ID. + * Used when loading purchases to enrich entitlements with transaction data. + */ + val entitlementsByProductId: Map> + get() = snapshot.entitlementsByProduct internal var activeDeviceEntitlements: Set - get() = _activeDeviceEntitlements + get() = snapshot.activeDeviceEntitlements set(value) { - _activeDeviceEntitlements.clear() - _activeDeviceEntitlements.addAll(value) + update(EntitlementsState.Updates.SetDeviceEntitlements(value)) } /** * All entitlements, regardless of whether they're active or not. + * Includes web entitlements from the latest redemption response. */ val all: Set - get() = _all.toSet() + entitlementsByProduct.values.flatten() + web.toSet() + get() = snapshot.all /** * The active entitlements. @@ -87,143 +90,61 @@ class Entitlements( * keeping the highest priority version of each and merging productIds. */ val active: Set - get() = mergeEntitlementsPrioritized((backingActive + _activeDeviceEntitlements + web).toList()).toSet() + get() = snapshot.active /** * The inactive entitlements. */ val inactive: Set - get() = _inactive.toSet() + all.minus(active) - - init { - try { - storage.read(StoredSubscriptionStatus)?.let { - setSubscriptionStatus(it) - } - } catch (e: ClassCastException) { - // Handle corrupted cache data - reset to Unknown status - storage.delete(StoredSubscriptionStatus) - setSubscriptionStatus(SubscriptionStatus.Unknown) - } - try { - storage.read(StoredEntitlementsByProductId)?.let { - entitlementsByProduct.putAll(it) - } - } catch (e: ClassCastException) { - // Handle corrupted cache data - storage.delete(StoredEntitlementsByProductId) - } - - scope.launch { - status.collect { - storage.write(StoredSubscriptionStatus, it) - } - } - } + get() = snapshot.inactive /** * Sets the entitlement status and updates the corresponding entitlement collections. + * + * The state update is synchronous; the new status is persisted right away. */ fun setSubscriptionStatus(value: SubscriptionStatus) { when (value) { is SubscriptionStatus.Active -> { if (value.entitlements.isEmpty()) { - setSubscriptionStatus(SubscriptionStatus.Inactive) + update(EntitlementsState.Updates.SetInactive) } else { - val entitlements = value.entitlements.toList().toSet() - backingActive.addAll(entitlements.filter { it.isActive }) - _all.addAll(entitlements) - _inactive.removeAll(entitlements) - _status.value = value + update(EntitlementsState.Updates.SetActive(value.entitlements.toSet())) } } - is SubscriptionStatus.Inactive -> { - _activeDeviceEntitlements.clear() - backingActive.clear() - _inactive.clear() - _status.value = value - } - - is SubscriptionStatus.Unknown -> { - backingActive.clear() - _activeDeviceEntitlements.clear() - _inactive.clear() - _status.value = value - } + is SubscriptionStatus.Inactive -> update(EntitlementsState.Updates.SetInactive) + is SubscriptionStatus.Unknown -> update(EntitlementsState.Updates.SetUnknown) } - } - - /** - * Returns a Set of Entitlements belonging to a given productId. - * - * @param id A String representing a productId - * @return A Set of Entitlements - */ - - private fun checkFor( - toCheck: List, - isExact: Boolean = true, - ): Set? { - if (toCheck.isEmpty()) return null - val item = toCheck.first() - val next = toCheck.drop(1) - return entitlementsByProduct.entries - .firstOrNull { - ( - if (isExact) { - it.key == item - } else { - it.key.contains(item) - } - ) && - it.value.isNotEmpty() - }?.value ?: checkFor(next, isExact) + _status.value = snapshot.status + persist(StoredSubscriptionStatus, snapshot.status) } /** * Checks for entitlements belonging to the product. * First checks exact matches, then checks containing matches - * by product ID + baseplan and productId so user doesn't remain without entitlements + * by product ID + baseplan and productId so user doesn't remain without entitlements * if they purchased the product. This ensures users dont lose access for their subscription. */ - internal fun byProductId(id: String): Set { - val decomposedProductIds = DecomposedProductIds.from(id) - return checkFor( - listOf( - decomposedProductIds.fullId, - "${decomposedProductIds.subscriptionId}:${decomposedProductIds.basePlanId ?: ""}:${decomposedProductIds.offerType.specificId ?: ""}", - "${decomposedProductIds.subscriptionId}:${decomposedProductIds.basePlanId ?: ""}", - ), - ) ?: checkFor( - listOf( - "${decomposedProductIds.subscriptionId}:${decomposedProductIds.basePlanId ?: ""}:", - decomposedProductIds.subscriptionId, - ), - isExact = false, - ) ?: emptySet() - } + internal fun byProductId(id: String): Set = snapshot.byProductId(id) /** * Returns a Set of Entitlements belonging to given product IDs. - * - * @param ids A Set of Strings representing product IDs - * @return A Set of Entitlements */ - fun byProductIds(ids: Set): Set = ids.flatMap { byProductId(it) }.toSet() + fun byProductIds(ids: Set): Set = snapshot.byProductIds(ids) + + /** + * Replaces the active web entitlements from a redemption response. + */ + internal fun setWebEntitlements(entitlements: Set) { + update(EntitlementsState.Updates.SetWebEntitlements(entitlements)) + } /** * Updates the entitlements associated with product IDs and persists them to storage. */ internal fun addEntitlementsByProductId(idToEntitlements: Map>) { - entitlementsByProduct.putAll( - idToEntitlements - .mapValues { (_, entitlements) -> - entitlements.toSet() - }.toMap(), - ) - _all.clear() - _all.addAll(entitlementsByProduct.values.flatten()) - storage.write(StoredEntitlementsByProductId, entitlementsByProduct) + update(EntitlementsState.Updates.AddProductEntitlements(idToEntitlements)) + persist(StoredEntitlementsByProductId, snapshot.entitlementsByProduct) } } diff --git a/superwall/src/main/java/com/superwall/sdk/store/EntitlementsContext.kt b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsContext.kt new file mode 100644 index 000000000..59ae4a63b --- /dev/null +++ b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsContext.kt @@ -0,0 +1,11 @@ +package com.superwall.sdk.store + +import com.superwall.sdk.misc.primitives.BaseContext + +/** + * All dependencies available to entitlements [EntitlementsState.Actions]. + * + * Actions see only [EntitlementsState] via [actor] plus the storage helpers + * inherited from [BaseContext]. + */ +interface EntitlementsContext : BaseContext diff --git a/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt new file mode 100644 index 000000000..7fd9fa8bb --- /dev/null +++ b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt @@ -0,0 +1,198 @@ +package com.superwall.sdk.store + +import com.superwall.sdk.billing.DecomposedProductIds +import com.superwall.sdk.misc.primitives.Reducer +import com.superwall.sdk.misc.primitives.TypedAction +import com.superwall.sdk.models.customer.mergeEntitlementsPrioritized +import com.superwall.sdk.models.entitlements.Entitlement +import com.superwall.sdk.models.entitlements.SubscriptionStatus +import com.superwall.sdk.storage.LatestRedemptionResponse +import com.superwall.sdk.storage.Storage +import com.superwall.sdk.storage.StoredEntitlementsByProductId +import com.superwall.sdk.storage.StoredSubscriptionStatus + +data class EntitlementsState( + val status: SubscriptionStatus = SubscriptionStatus.Unknown, + val entitlementsByProduct: Map> = emptyMap(), + val activeDeviceEntitlements: Set = emptySet(), + val backingActive: Set = emptySet(), + /** Active web entitlements from the latest redemption response. */ + val webEntitlements: Set = emptySet(), + /** Tracks all entitlements seen from status updates. */ + val allTracked: Set = emptySet(), +) { + // -- Derived properties -- + + val all: Set + get() = allTracked + entitlementsByProduct.values.flatten() + webEntitlements + + val active: Set + get() = + mergeEntitlementsPrioritized( + (backingActive + activeDeviceEntitlements + webEntitlements).toList(), + ).toSet() + + val inactive: Set + get() = all - active + + // -- Product ID lookup (pure, operates on current state) -- + + internal fun byProductId(id: String): Set { + val decomposed = DecomposedProductIds.from(id) + return checkFor( + listOf( + decomposed.fullId, + "${decomposed.subscriptionId}:${decomposed.basePlanId ?: ""}:${decomposed.offerType.specificId ?: ""}", + "${decomposed.subscriptionId}:${decomposed.basePlanId ?: ""}", + ), + ) ?: checkFor( + listOf( + "${decomposed.subscriptionId}:${decomposed.basePlanId ?: ""}:", + decomposed.subscriptionId, + ), + isExact = false, + ) ?: emptySet() + } + + fun byProductIds(ids: Set): Set = ids.flatMap { byProductId(it) }.toSet() + + private fun checkFor( + toCheck: List, + isExact: Boolean = true, + ): Set? { + if (toCheck.isEmpty()) return null + val item = toCheck.first() + val next = toCheck.drop(1) + return entitlementsByProduct.entries + .firstOrNull { + (if (isExact) it.key == item else it.key.contains(item)) && + it.value.isNotEmpty() + }?.value ?: checkFor(next, isExact) + } + + // ----------------------------------------------------------------------- + // Pure state mutations — (EntitlementsState) -> EntitlementsState + // ----------------------------------------------------------------------- + + internal sealed class Updates( + override val reduce: (EntitlementsState) -> EntitlementsState, + ) : Reducer { + data class SetActive( + val entitlements: Set, + ) : Updates({ state -> + state.copy( + status = SubscriptionStatus.Active(entitlements), + backingActive = state.backingActive + entitlements.filter { it.isActive }, + allTracked = state.allTracked + entitlements, + ) + }) + + object SetInactive : Updates({ state -> + state.copy( + status = SubscriptionStatus.Inactive, + activeDeviceEntitlements = emptySet(), + backingActive = emptySet(), + ) + }) + + object SetUnknown : Updates({ state -> + state.copy( + status = SubscriptionStatus.Unknown, + backingActive = emptySet(), + activeDeviceEntitlements = emptySet(), + ) + }) + + data class AddProductEntitlements( + val idToEntitlements: Map>, + ) : Updates({ state -> + val newProducts = + state.entitlementsByProduct + + idToEntitlements.mapValues { (_, v) -> v.toSet() } + state.copy(entitlementsByProduct = newProducts) + }) + + data class SetDeviceEntitlements( + val entitlements: Set, + ) : Updates({ state -> + state.copy(activeDeviceEntitlements = entitlements) + }) + + data class SetWebEntitlements( + val entitlements: Set, + ) : Updates({ state -> + state.copy(webEntitlements = entitlements) + }) + } + + // ----------------------------------------------------------------------- + // Actions — async work via EntitlementsContext + // ----------------------------------------------------------------------- + + internal sealed class Actions( + override val execute: suspend EntitlementsContext.() -> Unit, + ) : TypedAction +} + +/** + * Builds initial EntitlementsState from storage BEFORE the actor starts. + */ +internal fun createInitialEntitlementsState(storage: Storage): EntitlementsState { + val status = + try { + storage.read(StoredSubscriptionStatus) + } catch (e: ClassCastException) { + storage.delete(StoredSubscriptionStatus) + null + } + + val productEntitlements = + try { + storage.read(StoredEntitlementsByProductId) + } catch (e: ClassCastException) { + storage.delete(StoredEntitlementsByProductId) + null + } + + var state = EntitlementsState() + + // Restore product entitlements, then the status. allTracked only holds + // status entitlements; `all` adds the product ones from entitlementsByProduct. + if (productEntitlements != null) { + state = EntitlementsState.Updates.AddProductEntitlements(productEntitlements).reduce(state) + } + + // Replay status to populate backingActive/allTracked correctly + if (status != null) { + state = + when (status) { + is SubscriptionStatus.Active -> { + if (status.entitlements.isEmpty()) { + EntitlementsState.Updates.SetInactive.reduce(state) + } else { + EntitlementsState.Updates.SetActive(status.entitlements.toSet()).reduce(state) + } + } + is SubscriptionStatus.Inactive -> EntitlementsState.Updates.SetInactive.reduce(state) + is SubscriptionStatus.Unknown -> state + } + } + + // Restore web entitlements from latest redemption response + val webEntitlements = + try { + storage + .read(LatestRedemptionResponse) + ?.customerInfo + ?.entitlements + ?.filter { it.isActive } + ?.toSet() + } catch (_: Exception) { + null + } + if (!webEntitlements.isNullOrEmpty()) { + state = EntitlementsState.Updates.SetWebEntitlements(webEntitlements).reduce(state) + } + + return state +} diff --git a/superwall/src/main/java/com/superwall/sdk/store/testmode/TestMode.kt b/superwall/src/main/java/com/superwall/sdk/store/testmode/TestMode.kt index cef1689fa..47c914029 100644 --- a/superwall/src/main/java/com/superwall/sdk/store/testmode/TestMode.kt +++ b/superwall/src/main/java/com/superwall/sdk/store/testmode/TestMode.kt @@ -1,16 +1,17 @@ package com.superwall.sdk.store.testmode -import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent -import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent.TestModeModal.State +import android.app.Activity +import com.superwall.sdk.analytics.internal.trackable.TrackableSuperwallEvent import com.superwall.sdk.logger.LogLevel import com.superwall.sdk.logger.LogScope import com.superwall.sdk.logger.Logger import com.superwall.sdk.misc.ActivityProvider import com.superwall.sdk.misc.CurrentActivityTracker import com.superwall.sdk.misc.Either -import com.superwall.sdk.misc.fold +import com.superwall.sdk.misc.IOScope +import com.superwall.sdk.misc.primitives.SequentialActor +import com.superwall.sdk.misc.primitives.StateActor import com.superwall.sdk.models.config.Config -import com.superwall.sdk.models.entitlements.Entitlement import com.superwall.sdk.models.entitlements.SubscriptionStatus import com.superwall.sdk.network.NetworkError import com.superwall.sdk.storage.IsTestModeActiveSubscription @@ -21,42 +22,55 @@ import com.superwall.sdk.store.Entitlements import com.superwall.sdk.store.abstractions.product.StoreProduct import com.superwall.sdk.store.testmode.models.SuperwallEntitlementRef import com.superwall.sdk.store.testmode.models.SuperwallProduct -import com.superwall.sdk.store.testmode.models.SuperwallProductPlatform import com.superwall.sdk.store.testmode.models.SuperwallProductsResponse -import com.superwall.sdk.store.testmode.models.TestStoreUserType import com.superwall.sdk.store.testmode.ui.EntitlementSelection import com.superwall.sdk.store.testmode.ui.EntitlementStateOption import com.superwall.sdk.store.testmode.ui.TestModeModal +import com.superwall.sdk.store.testmode.ui.TestModeModalResult +import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.withTimeoutOrNull import kotlin.time.Duration import kotlin.time.Duration.Companion.seconds /** - * The single test-mode surface: holds the activation state (products, - * entitlement selections, settings persistence) AND runs the activation UI - * flow (`activate` → refresh products → present modal). + * Test-mode manager. * - * Not exactly a "manager" — the UI flow pieces (activity lookup, subscription - * products fetch, modal presentation) are injected as thin lambdas so this - * class stays testable and config-slice-free. + * Implements [TestModeContext] directly so [TestModeState.Actions] receive + * `this` as their receiver — same pattern as + * [com.superwall.sdk.identity.IdentityManager] / [com.superwall.sdk.identity.IdentityContext]. + * + * State is held in a [SequentialActor]; pure mutations go through + * `update(Updates.X)` (CAS-atomic), async work (network + modal) is + * dispatched as [TestModeState.Actions]. */ class TestMode( - private val storage: Storage, - private val isTestEnvironment: Boolean = Companion.isTestEnvironment, - // Activation UI hooks — all default to no-ops so unit tests exercising - // state management can construct `TestMode(storage)` without wiring the - // whole UI/network surface. - private val getSuperwallProducts: suspend () -> Either = { + override val storage: Storage, + override val isTestEnvironment: Boolean = Companion.isTestEnvironment, + override val getSuperwallProducts: suspend () -> Either = { Either.Failure(NetworkError.Unknown()) }, - private val entitlements: Entitlements? = null, - private val activityProvider: () -> ActivityProvider? = { null }, - private val activityTracker: () -> CurrentActivityTracker? = { null }, - private val hasExternalPurchaseController: () -> Boolean = { false }, - private val apiKey: () -> String = { "" }, - private val dashboardBaseUrl: () -> String = { "" }, - private val track: suspend (InternalSuperwallEvent) -> Unit = { }, -) { + override val entitlements: Entitlements? = null, + override val activityProvider: () -> ActivityProvider? = { null }, + override val activityTracker: () -> CurrentActivityTracker? = { null }, + override val hasExternalPurchaseController: () -> Boolean = { false }, + override val apiKey: () -> String = { "" }, + override val dashboardBaseUrl: () -> String = { "" }, + override val tracker: suspend (TrackableSuperwallEvent) -> Unit = { }, + override val showModal: suspend ( + activity: Activity, + reason: String, + hasPurchaseController: Boolean, + availableEntitlements: List, + apiKey: String, + dashboardBaseUrl: String, + savedSettings: TestModeSettings?, + ) -> TestModeModalResult = { activity, reason, hasPC, available, ak, db, saved -> + TestModeModal.show(activity, reason, hasPC, available, ak, db, saved) + }, + private val ioScope: CoroutineScope = IOScope(), + override val actor: StateActor = + SequentialActor(TestModeState.Inactive, ioScope), +) : TestModeContext { companion object { val isTestEnvironment: Boolean by lazy { try { @@ -78,15 +92,14 @@ class TestMode( } } - var state: TestModeState = TestModeState.Inactive - private set + override val scope: CoroutineScope get() = ioScope + + // ---- Read accessors (snapshot of state.value) ------------------------- - // Convenience accessors - val isTestMode: Boolean get() = state is TestModeState.Active - val testModeReason: TestModeReason? get() = (state as? TestModeState.Active)?.reason - private val session: TestModeSessionData? get() = (state as? TestModeState.Active)?.session + val isTestMode: Boolean get() = state.value is TestModeState.Active + val testModeReason: TestModeReason? get() = (state.value as? TestModeState.Active)?.reason + private val session: TestModeSessionData? get() = state.value.sessionOrNull - // Backward-compatible session data accessors (return sensible defaults when inactive) val products: List get() = session?.products ?: emptyList() internal val testProductsByFullId: Map get() = session?.testProductsByFullId ?: emptyMap() val testEntitlementIds: Set get() = session?.entitlementIds ?: emptySet() @@ -94,6 +107,8 @@ class TestMode( val freeTrialOverride: FreeTrialOverride get() = session?.freeTrialOverride ?: FreeTrialOverride.UseDefault val overriddenSubscriptionStatus: SubscriptionStatus? get() = session?.overriddenSubscriptionStatus + // ---- Pure-state mutators (synchronous, CAS-atomic) -------------------- + fun evaluateTestMode( config: Config, bundleId: String, @@ -101,186 +116,118 @@ class TestMode( aliasId: String?, testModeBehavior: TestModeBehavior = TestModeBehavior.AUTOMATIC, ) { - when (testModeBehavior) { - TestModeBehavior.NEVER -> { - deactivateIfActive() - return - } - - TestModeBehavior.ALWAYS -> { - activateWithReason(TestModeReason.TestModeOption) - return - } - - TestModeBehavior.WHEN_ENABLED_FOR_USER -> { - if (checkConfigMatch(config, appUserId, aliasId)) return - deactivateIfActive() - return - } - - TestModeBehavior.AUTOMATIC -> { - // Skip in test environments (JUnit on classpath) - if (isTestEnvironment) { - deactivateIfActive() - return - } - if (checkConfigMatch(config, appUserId, aliasId)) return - if (checkPackageNameMismatch(config, bundleId)) return - deactivateIfActive() - } + val newReason = + TestModeLogic.evaluate( + config = config, + bundleId = bundleId, + appUserId = appUserId, + aliasId = aliasId, + behavior = testModeBehavior, + isTestEnvironment = isTestEnvironment, + ) + if (newReason == null) { + if (isTestMode) clearTestModeState() + return } - } - - private fun deactivateIfActive() { - if (isTestMode) { - clearTestModeState() + val previousReason = testModeReason + update(TestModeState.Updates.SetActive(newReason)) + if (previousReason != null && previousReason != newReason) { + storage.write(IsTestModeActiveSubscription, false) } - } - - private fun activateWithReason(reason: TestModeReason) { - val current = state - state = - if (current is TestModeState.Active) { - if (current.reason != reason) { - current.session.entitlementIds.clear() - current.session.entitlementSelections = emptyList() - current.session.overriddenSubscriptionStatus = null - storage.write(IsTestModeActiveSubscription, false) - } - current.copy(reason = reason) - } else { - TestModeState.Active(reason = reason) - } Logger.debug( LogLevel.info, LogScope.superwallCore, - "Test mode activated: ${testModeReason?.description}", + "Test mode activated: ${newReason.description}", ) } - suspend fun awaitTestProducts(timeout: Duration = 5.seconds) { - val s = session ?: return - withTimeoutOrNull(timeout) { s.productsLoaded.await() } - } - - private fun checkConfigMatch( - config: Config, - appUserId: String?, - aliasId: String?, - ): Boolean { - val testUsers = config.testModeUserIds ?: return false - for (testUser in testUsers) { - val match = - when (testUser.type) { - TestStoreUserType.UserId -> appUserId == testUser.value - TestStoreUserType.AliasId -> aliasId == testUser.value - } - if (match) { - activateWithReason(TestModeReason.ConfigMatch(matchedId = testUser.value)) - return true - } - } - return false + fun setProducts(products: List) { + update(TestModeState.Updates.UpdateSession { it.copy(products = products) }) } - private fun checkPackageNameMismatch( - config: Config, - actualPackageName: String, - ): Boolean { - val expectedPackageName = config.bundleIdConfig - if (expectedPackageName.isNullOrEmpty()) return false - if (expectedPackageName == actualPackageName) return false - // Treat as extension if actual starts with expected + "." - if (actualPackageName.startsWith("$expectedPackageName.")) return false - - activateWithReason( - TestModeReason.ApplicationIdMismatch( - expected = expectedPackageName, - actual = actualPackageName, - ), + fun setTestProducts(productsByFullId: Map) { + update( + TestModeState.Updates.UpdateSession { + it.copy(testProductsByFullId = productsByFullId) + }, ) - return true + session?.productsLoaded?.complete(Unit) } - fun setProducts(products: List) { - session?.products = products - } - - fun setTestProducts(productsByFullId: Map) { - session?.let { - it.testProductsByFullId = productsByFullId - it.productsLoaded.complete(Unit) - } + /** Suspend until the test product catalog has been loaded (or [timeout] elapses). No-op when inactive. */ + suspend fun awaitTestProducts(timeout: Duration = 5.seconds) { + val s = session ?: return + withTimeoutOrNull(timeout) { s.productsLoaded.await() } } fun fakePurchase(entitlementRefs: List) { - val ids = entitlementRefs.map { it.identifier } - session?.entitlementIds?.addAll(ids) + val ids = entitlementRefs.map { it.identifier }.toSet() + update( + TestModeState.Updates.UpdateSession { + it.copy(entitlementIds = it.entitlementIds + ids) + }, + ) storage.write(IsTestModeActiveSubscription, testEntitlementIds.isNotEmpty()) } fun setEntitlements(selections: List) { - val s = session ?: return - s.entitlementSelections = selections - s.entitlementIds.clear() - s.entitlementIds.addAll( - selections.filter { it.state.isActive }.map { it.identifier }, + val newIds = + selections.filter { it.state.isActive }.map { it.identifier }.toSet() + update( + TestModeState.Updates.UpdateSession { + it.copy(entitlementSelections = selections, entitlementIds = newIds) + }, ) - storage.write(IsTestModeActiveSubscription, s.entitlementIds.isNotEmpty()) + storage.write(IsTestModeActiveSubscription, newIds.isNotEmpty()) } - fun setEntitlements(ids: Set) { + fun setEntitlements(ids: Set) = setEntitlements( - ids.map { EntitlementSelection(identifier = it, state = EntitlementStateOption.Subscribed) }, + ids.map { + EntitlementSelection(identifier = it, state = EntitlementStateOption.Subscribed) + }, ) - } fun resetEntitlements() { - session?.entitlementIds?.clear() - session?.entitlementSelections = emptyList() + update( + TestModeState.Updates.UpdateSession { + it.copy(entitlementIds = emptySet(), entitlementSelections = emptyList()) + }, + ) storage.write(IsTestModeActiveSubscription, false) } fun setFreeTrialOverride(override: FreeTrialOverride) { - session?.freeTrialOverride = override + update(TestModeState.Updates.UpdateSession { it.copy(freeTrialOverride = override) }) } - fun shouldShowFreeTrial(hasFreeTrial: Boolean): Boolean = - when (freeTrialOverride) { - FreeTrialOverride.UseDefault -> hasFreeTrial - FreeTrialOverride.ForceAvailable -> true - FreeTrialOverride.ForceUnavailable -> false - } + fun setOverriddenSubscriptionStatus(status: SubscriptionStatus?) { + update(TestModeState.Updates.UpdateSession { it.copy(overriddenSubscriptionStatus = status) }) + } fun clearTestModeState() { - state = TestModeState.Inactive + update(TestModeState.Updates.SetInactive) storage.delete(IsTestModeActiveSubscription) clearSettings() } - fun buildSubscriptionStatus(): SubscriptionStatus { - if (testEntitlementIds.isEmpty()) { - return SubscriptionStatus.Inactive - } - val activeSelections = testEntitlementSelections.filter { it.state.isActive } - return if (activeSelections.isNotEmpty()) { - SubscriptionStatus.Active( - activeSelections.map { it.toEntitlement() }.toSet(), - ) - } else { - SubscriptionStatus.Active( - testEntitlementIds.map { Entitlement(it) }.toSet(), - ) + // ---- Derived helpers -------------------------------------------------- + + fun shouldShowFreeTrial(hasFreeTrial: Boolean): Boolean = + when (freeTrialOverride) { + FreeTrialOverride.UseDefault -> hasFreeTrial + FreeTrialOverride.ForceAvailable -> true + FreeTrialOverride.ForceUnavailable -> false } - } - fun setOverriddenSubscriptionStatus(status: SubscriptionStatus?) { - session?.overriddenSubscriptionStatus = status - } + fun buildSubscriptionStatus(): SubscriptionStatus = buildSubscriptionStatus(state.value) fun entitlementsForProduct(product: SuperwallProduct): List = product.entitlements - fun allEntitlements(): Set = products.flatMap { it.entitlements.map { e -> e.identifier } }.toSet() + fun allEntitlements(): Set = + products.flatMap { it.entitlements.map { e -> e.identifier } }.toSet() + + // ---- Settings persistence -------------------------------------------- fun saveSettings() { val settings = @@ -297,106 +244,18 @@ class TestMode( storage.delete(StoredTestModeSettings) } - // ---- Activation UI flow ------------------------------------------------ + // ---- Async activation flow ------------------------------------------- /** - * Refresh the test product catalog and (when [justActivated] is true) - * present the test-mode modal. Must be called off the actor queue — - * [presentModal] blocks on user interaction. + * Refresh the test product catalog and (when [justActivated]) present + * the modal. Runs as a [TestModeState.Actions.Activate] action and suspends + * until it completes, so callers that must not wait on the modal's blocking + * UI (e.g. ConfigState) launch it in their own scope. */ suspend fun activate( config: Config, justActivated: Boolean, ) { - refreshProducts() - if (justActivated) { - presentModal(config) - } - } - - private suspend fun refreshProducts() { - try { - getSuperwallProducts().fold( - onSuccess = { response -> - val androidProducts = - response.data.filter { - it.platform == SuperwallProductPlatform.ANDROID && it.price != null - } - setProducts(androidProducts) - - val productsByFullId = - androidProducts.associate { superwallProduct -> - val testProduct = TestStoreProduct(superwallProduct) - superwallProduct.identifier to StoreProduct(testProduct) - } - setTestProducts(productsByFullId) - - Logger.debug( - LogLevel.info, - LogScope.superwallCore, - "Test mode: loaded ${androidProducts.size} products", - ) - }, - onFailure = { error -> - Logger.debug( - LogLevel.error, - LogScope.superwallCore, - "Test mode: failed to fetch products - ${error.message}", - ) - }, - ) - } finally { - session?.productsLoaded?.complete(Unit) - } - } - - private suspend fun presentModal(config: Config) { - val activity = - activityTracker()?.getCurrentActivity() - ?: activityProvider()?.getCurrentActivity() - ?: activityTracker()?.awaitActivity(10.seconds) - if (activity == null) { - Logger.debug( - LogLevel.warn, - LogScope.superwallCore, - "Test mode modal could not be presented: no activity available. Setting default subscription status.", - ) - val status = buildSubscriptionStatus() - setOverriddenSubscriptionStatus(status) - entitlements?.setSubscriptionStatus(status) - return - } - - track(InternalSuperwallEvent.TestModeModal(State.Open)) - - val reason = testModeReason?.description ?: "Test mode activated" - val allEntitlements = - config.productsV3 - ?.flatMap { it.entitlements.map { e -> e.id } } - ?.distinct() - ?.sorted() - ?: emptyList() - - val savedSettings = loadSettings() - - val result = - TestModeModal.show( - activity = activity, - reason = reason, - hasPurchaseController = hasExternalPurchaseController(), - availableEntitlements = allEntitlements, - apiKey = apiKey(), - dashboardBaseUrl = dashboardBaseUrl(), - savedSettings = savedSettings, - ) - - setFreeTrialOverride(result.freeTrialOverride) - setEntitlements(result.entitlements) - saveSettings() - val status = buildSubscriptionStatus() - setOverriddenSubscriptionStatus(status) - entitlements?.setSubscriptionStatus(status) - - track(InternalSuperwallEvent.TestModeModal(State.Close)) + immediate(TestModeState.Actions.Activate(config, justActivated)) } } diff --git a/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeContext.kt b/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeContext.kt new file mode 100644 index 000000000..3778456c7 --- /dev/null +++ b/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeContext.kt @@ -0,0 +1,39 @@ +package com.superwall.sdk.store.testmode + +import android.app.Activity +import com.superwall.sdk.misc.ActivityProvider +import com.superwall.sdk.misc.CurrentActivityTracker +import com.superwall.sdk.misc.Either +import com.superwall.sdk.misc.primitives.BaseContext +import com.superwall.sdk.network.NetworkError +import com.superwall.sdk.storage.TestModeSettings +import com.superwall.sdk.store.Entitlements +import com.superwall.sdk.store.testmode.models.SuperwallProductsResponse +import com.superwall.sdk.store.testmode.ui.TestModeModalResult + +/** + * Dependencies available to [TestModeState.Actions]. + * + * Implemented directly by [TestMode] — actions receive the manager itself + * as their context, mirroring the [com.superwall.sdk.identity.IdentityManager] + * / [com.superwall.sdk.identity.IdentityContext] pattern. + */ +interface TestModeContext : BaseContext { + val isTestEnvironment: Boolean + val entitlements: Entitlements? + val getSuperwallProducts: suspend () -> Either + val activityProvider: () -> ActivityProvider? + val activityTracker: () -> CurrentActivityTracker? + val hasExternalPurchaseController: () -> Boolean + val apiKey: () -> String + val dashboardBaseUrl: () -> String + val showModal: suspend ( + activity: Activity, + reason: String, + hasPurchaseController: Boolean, + availableEntitlements: List, + apiKey: String, + dashboardBaseUrl: String, + savedSettings: TestModeSettings?, + ) -> TestModeModalResult +} diff --git a/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeLogic.kt b/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeLogic.kt new file mode 100644 index 000000000..9c0aa28f4 --- /dev/null +++ b/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeLogic.kt @@ -0,0 +1,64 @@ +package com.superwall.sdk.store.testmode + +import com.superwall.sdk.models.config.Config +import com.superwall.sdk.store.testmode.models.TestStoreUserType + +/** Pure decision logic — `null` means deactivate. */ +internal object TestModeLogic { + fun evaluate( + config: Config, + bundleId: String, + appUserId: String?, + aliasId: String?, + behavior: TestModeBehavior, + isTestEnvironment: Boolean, + ): TestModeReason? = + when (behavior) { + TestModeBehavior.NEVER -> null + TestModeBehavior.ALWAYS -> TestModeReason.TestModeOption + TestModeBehavior.WHEN_ENABLED_FOR_USER -> + checkConfigMatch(config, appUserId, aliasId) + TestModeBehavior.AUTOMATIC -> { + if (isTestEnvironment) { + null + } else { + checkConfigMatch(config, appUserId, aliasId) + ?: checkPackageNameMismatch(config, bundleId) + } + } + } + + private fun checkConfigMatch( + config: Config, + appUserId: String?, + aliasId: String?, + ): TestModeReason? { + val testUsers = config.testModeUserIds ?: return null + for (testUser in testUsers) { + val match = + when (testUser.type) { + TestStoreUserType.UserId -> appUserId == testUser.value + TestStoreUserType.AliasId -> aliasId == testUser.value + } + if (match) { + return TestModeReason.ConfigMatch(matchedId = testUser.value) + } + } + return null + } + + private fun checkPackageNameMismatch( + config: Config, + actualPackageName: String, + ): TestModeReason? { + val expectedPackageName = config.bundleIdConfig + if (expectedPackageName.isNullOrEmpty()) return null + if (expectedPackageName == actualPackageName) return null + // Treat actual = expected + ".something" as an extension/variant — not a mismatch. + if (actualPackageName.startsWith("$expectedPackageName.")) return null + return TestModeReason.ApplicationIdMismatch( + expected = expectedPackageName, + actual = actualPackageName, + ) + } +} diff --git a/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeState.kt b/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeState.kt index 89aadbbe3..6db49ad7e 100644 --- a/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeState.kt +++ b/superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeState.kt @@ -1,10 +1,24 @@ package com.superwall.sdk.store.testmode +import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent +import com.superwall.sdk.logger.LogLevel +import com.superwall.sdk.logger.LogScope +import com.superwall.sdk.logger.Logger +import com.superwall.sdk.misc.fold +import com.superwall.sdk.misc.primitives.Reducer +import com.superwall.sdk.misc.primitives.TypedAction +import com.superwall.sdk.models.config.Config +import com.superwall.sdk.models.entitlements.Entitlement import com.superwall.sdk.models.entitlements.SubscriptionStatus +import com.superwall.sdk.storage.IsTestModeActiveSubscription +import com.superwall.sdk.storage.StoredTestModeSettings +import com.superwall.sdk.storage.TestModeSettings import com.superwall.sdk.store.abstractions.product.StoreProduct import com.superwall.sdk.store.testmode.models.SuperwallProduct +import com.superwall.sdk.store.testmode.models.SuperwallProductPlatform import com.superwall.sdk.store.testmode.ui.EntitlementSelection import kotlinx.coroutines.CompletableDeferred +import kotlin.time.Duration.Companion.seconds sealed class TestModeState { data object Inactive : TestModeState() @@ -13,15 +27,199 @@ sealed class TestModeState { val reason: TestModeReason, val session: TestModeSessionData = TestModeSessionData(), ) : TestModeState() + + /** Per-activation working set. Immutable — every mutation produces a copy. */ + data class TestModeSessionData( + val products: List = emptyList(), + val testProductsByFullId: Map = emptyMap(), + val entitlementIds: Set = emptySet(), + val entitlementSelections: List = emptyList(), + val freeTrialOverride: FreeTrialOverride = FreeTrialOverride.UseDefault, + val overriddenSubscriptionStatus: SubscriptionStatus? = null, + /** Completed once the product catalog refresh finishes (success or failure). Carried through `copy()`. */ + val productsLoaded: CompletableDeferred = CompletableDeferred(), + ) + + val sessionOrNull: TestModeSessionData? get() = (this as? Active)?.session + + internal sealed class Updates( + override val reduce: (TestModeState) -> TestModeState, + ) : Reducer { + /** + * Activate with [reason]. An already-active session is preserved (products, + * free-trial override); when the reason changes, entitlement selections and + * the overridden status are cleared so the new reason starts clean. + */ + data class SetActive(val reason: TestModeReason) : Updates({ state -> + when { + state !is Active -> Active(reason) + state.reason == reason -> state + else -> + state.copy( + reason = reason, + session = + state.session.copy( + entitlementIds = emptySet(), + entitlementSelections = emptyList(), + overriddenSubscriptionStatus = null, + ), + ) + } + }) + + object SetInactive : Updates({ Inactive }) + + /** Mutate the active session. No-op when state is Inactive. */ + data class UpdateSession( + val transform: (TestModeSessionData) -> TestModeSessionData, + ) : Updates({ state -> + when (state) { + is Active -> state.copy(session = transform(state.session)) + Inactive -> state + } + }) + } + + internal sealed class Actions( + override val execute: suspend TestModeContext.() -> Unit, + ) : TypedAction { + /** Refresh test-product catalog from the network. */ + object RefreshProducts : Actions({ + try { + getSuperwallProducts().fold( + onSuccess = { response -> + val androidProducts = + response.data.filter { + it.platform == SuperwallProductPlatform.ANDROID && it.price != null + } + val productsByFullId = + androidProducts.associate { superwallProduct -> + val testProduct = TestStoreProduct(superwallProduct) + superwallProduct.identifier to StoreProduct(testProduct) + } + update( + Updates.UpdateSession { + it.copy( + products = androidProducts, + testProductsByFullId = productsByFullId, + ) + }, + ) + Logger.debug( + LogLevel.info, + LogScope.superwallCore, + "Test mode: loaded ${androidProducts.size} products", + ) + }, + onFailure = { error -> + Logger.debug( + LogLevel.error, + LogScope.superwallCore, + "Test mode: failed to fetch products - ${error.message}", + ) + }, + ) + } finally { + state.value.sessionOrNull?.productsLoaded?.complete(Unit) + } + }) + + /** Refresh products + (when newly activated) present the modal. */ + data class Activate( + val config: Config, + val justActivated: Boolean, + ) : Actions(exec@{ + immediate(RefreshProducts) + + if (!justActivated) return@exec + + // ---- Present modal ------------------------------------------ + val activity = + activityTracker()?.getCurrentActivity() + ?: activityProvider()?.getCurrentActivity() + ?: activityTracker()?.awaitActivity(10.seconds) + + if (activity == null) { + Logger.debug( + LogLevel.warn, + LogScope.superwallCore, + "Test mode modal could not be presented: no activity available. Setting default subscription status.", + ) + val status = buildSubscriptionStatus(state.value) + update(Updates.UpdateSession { it.copy(overriddenSubscriptionStatus = status) }) + entitlements?.setSubscriptionStatus(status) + return@exec + } + + track(InternalSuperwallEvent.TestModeModal(InternalSuperwallEvent.TestModeModal.State.Open)) + + val reason = + (state.value as? Active)?.reason?.description ?: "Test mode activated" + val allEntitlements = + config.productsV3 + ?.flatMap { it.entitlements.map { e -> e.id } } + ?.distinct() + ?.sorted() + ?: emptyList() + + val savedSettings = storage.read(StoredTestModeSettings) + + val result = + showModal( + activity, + reason, + hasExternalPurchaseController(), + allEntitlements, + apiKey(), + dashboardBaseUrl(), + savedSettings, + ) + + val newSelections = result.entitlements + val newIds = + newSelections + .filter { it.state.isActive } + .map { it.identifier } + .toSet() + + update( + Updates.UpdateSession { session -> + session.copy( + freeTrialOverride = result.freeTrialOverride, + entitlementSelections = newSelections, + entitlementIds = newIds, + ) + }, + ) + storage.write(IsTestModeActiveSubscription, newIds.isNotEmpty()) + storage.write( + StoredTestModeSettings, + TestModeSettings( + entitlementSelections = newSelections, + freeTrialOverride = result.freeTrialOverride, + ), + ) + + val status = buildSubscriptionStatus(state.value) + update(Updates.UpdateSession { it.copy(overriddenSubscriptionStatus = status) }) + entitlements?.setSubscriptionStatus(status) + + track(InternalSuperwallEvent.TestModeModal(InternalSuperwallEvent.TestModeModal.State.Close)) + }) + } } -class TestModeSessionData { - var products: List = emptyList() - var testProductsByFullId: Map = emptyMap() +// Convenience alias — used widely from external call sites. +typealias TestModeSessionData = TestModeState.TestModeSessionData - val productsLoaded: CompletableDeferred = CompletableDeferred() - var entitlementIds: MutableSet = mutableSetOf() - var entitlementSelections: List = emptyList() - var freeTrialOverride: FreeTrialOverride = FreeTrialOverride.UseDefault - var overriddenSubscriptionStatus: SubscriptionStatus? = null +/** Pure derivation of subscription status from a [TestModeState] snapshot. */ +internal fun buildSubscriptionStatus(state: TestModeState): SubscriptionStatus { + val session = state.sessionOrNull ?: return SubscriptionStatus.Inactive + if (session.entitlementIds.isEmpty()) return SubscriptionStatus.Inactive + val activeSelections = session.entitlementSelections.filter { it.state.isActive } + return if (activeSelections.isNotEmpty()) { + SubscriptionStatus.Active(activeSelections.map { it.toEntitlement() }.toSet()) + } else { + SubscriptionStatus.Active(session.entitlementIds.map { Entitlement(it) }.toSet()) + } } diff --git a/superwall/src/main/java/com/superwall/sdk/store/testmode/ui/TestModeModal.kt b/superwall/src/main/java/com/superwall/sdk/store/testmode/ui/TestModeModal.kt index 4baedcb0e..c4cd562c9 100644 --- a/superwall/src/main/java/com/superwall/sdk/store/testmode/ui/TestModeModal.kt +++ b/superwall/src/main/java/com/superwall/sdk/store/testmode/ui/TestModeModal.kt @@ -29,7 +29,7 @@ import kotlinx.coroutines.launch import kotlinx.coroutines.withContext import kotlinx.serialization.json.Json -internal data class TestModeModalResult( +data class TestModeModalResult( val entitlements: List, val freeTrialOverride: FreeTrialOverride, ) diff --git a/superwall/src/main/java/com/superwall/sdk/web/WebPaywallRedeemer.kt b/superwall/src/main/java/com/superwall/sdk/web/WebPaywallRedeemer.kt index d9b5d6e1b..db5aa250d 100644 --- a/superwall/src/main/java/com/superwall/sdk/web/WebPaywallRedeemer.kt +++ b/superwall/src/main/java/com/superwall/sdk/web/WebPaywallRedeemer.kt @@ -80,6 +80,8 @@ class WebPaywallRedeemer( fun internallySetSubscriptionStatus(status: SubscriptionStatus) + fun setWebEntitlements(entitlements: Set) + suspend fun isPaywallVisible(): Boolean suspend fun triggerRestoreInPaywall() @@ -235,6 +237,12 @@ class WebPaywallRedeemer( ).fold( onSuccess = { storage.write(LatestRedemptionResponse, it) + factory.setWebEntitlements( + it.customerInfo + ?.entitlements + ?.filter { it.isActive } + ?.toSet() ?: emptySet(), + ) track( Redemptions( RedemptionState.Complete, @@ -465,6 +473,14 @@ class WebPaywallRedeemer( // Get active entitlements that remain after removing web sources or ones from the web if (withUserCodesRemoved != null) { storage.write(LatestRedemptionResponse, withUserCodesRemoved) + factory.setWebEntitlements( + withUserCodesRemoved.customerInfo + ?.entitlements + ?.filter { it.isActive } + ?.toSet() ?: emptySet(), + ) + } else { + factory.setWebEntitlements(emptySet()) } factory.internallySetSubscriptionStatus( SubscriptionStatus.Active( @@ -522,6 +538,11 @@ class WebPaywallRedeemer( LatestRedemptionResponse, updatedResponse, ) + // Publish only what was persisted, so the cached web + // entitlements always match what a cold start restores. + factory.setWebEntitlements( + newEntitlements.filter { it.isActive }.toSet(), + ) } // Trigger CustomerInfo merge diff --git a/superwall/src/test/java/com/superwall/sdk/SdkContextImplTest.kt b/superwall/src/test/java/com/superwall/sdk/SdkContextImplTest.kt index d5c2ae1bd..8bbbc2d7e 100644 --- a/superwall/src/test/java/com/superwall/sdk/SdkContextImplTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/SdkContextImplTest.kt @@ -1,7 +1,7 @@ package com.superwall.sdk import com.superwall.sdk.config.ConfigManager -import com.superwall.sdk.config.models.ConfigState +import com.superwall.sdk.config.ConfigState import com.superwall.sdk.models.config.Config import io.mockk.Runs import io.mockk.coEvery diff --git a/superwall/src/test/java/com/superwall/sdk/SubscriptionStatusDelegateOverrideTest.kt b/superwall/src/test/java/com/superwall/sdk/SubscriptionStatusDelegateOverrideTest.kt new file mode 100644 index 000000000..6afe1655f --- /dev/null +++ b/superwall/src/test/java/com/superwall/sdk/SubscriptionStatusDelegateOverrideTest.kt @@ -0,0 +1,310 @@ +package com.superwall.sdk + +import android.content.Context +import com.superwall.sdk.delegate.SuperwallDelegate +import com.superwall.sdk.delegate.SuperwallDelegateAdapter +import com.superwall.sdk.dependencies.DependencyContainer +import com.superwall.sdk.misc.IOScope +import com.superwall.sdk.models.entitlements.Entitlement +import com.superwall.sdk.models.entitlements.SubscriptionStatus +import com.superwall.sdk.storage.LocalStorage +import com.superwall.sdk.storage.Storable +import com.superwall.sdk.storage.StoredSubscriptionStatus +import com.superwall.sdk.store.makeEntitlements +import io.mockk.every +import io.mockk.mockk +import io.mockk.verify +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.cancel +import kotlinx.coroutines.flow.MutableStateFlow +import org.junit.After +import org.junit.Assert.assertEquals +import org.junit.Assert.assertTrue +import org.junit.Before +import org.junit.Test +import java.util.concurrent.CopyOnWriteArrayList + +/** + * Pins a pattern integrators rely on: overriding + * [SuperwallDelegate.subscriptionStatusDidChange] and calling + * [Superwall.setSubscriptionStatus] from inside it to add or remove + * entitlements before the rest of the app sees the status. + * + * It is not a documented API, but apps depend on it, so these tests run the + * real status listener from [Superwall] against real [com.superwall.sdk.store.Entitlements] + * and a real [SuperwallDelegateAdapter]. Only the dependency container around + * them is mocked. + */ +class SubscriptionStatusDelegateOverrideTest { + private lateinit var superwall: Superwall + private lateinit var storage: LocalStorage + private lateinit var ioScope: IOScope + private val calls = CopyOnWriteArrayList?, Set?>>() + + private val pro = Entitlement("pro") + private val legacy = Entitlement("legacy") + private val custom = Entitlement("custom") + + @Before + fun setUp() { + hasInitialized().value = true + ioScope = IOScope(Dispatchers.Unconfined) + storage = mockk(relaxed = true) + every { storage.read(any>()) } returns null + + val container = mockk(relaxed = true) + every { container.storage } returns storage + every { container.ioScope() } returns ioScope + every { container.delegateAdapter } returns SuperwallDelegateAdapter() + every { container.entitlements } returns makeEntitlements(storage, ioScope) + every { container.testMode.isTestMode } returns false + + superwall = + Superwall( + context = mockk(relaxed = true), + apiKey = "test", + purchaseController = null, + options = null, + activityProvider = null, + completion = null, + ) + Superwall::class.java.getDeclaredField("_dependencyContainer").apply { + isAccessible = true + set(superwall, container) + } + Superwall::class.java.getDeclaredMethod("addListeners").apply { + isAccessible = true + invoke(superwall) + } + } + + @After + fun tearDown() { + ioScope.cancel() + hasInitialized().value = false + } + + @Suppress("UNCHECKED_CAST") + private fun hasInitialized(): MutableStateFlow = + Superwall::class.java + .getDeclaredField("_hasInitialized") + .apply { isAccessible = true } + .get(null) as MutableStateFlow + + /** Installs a delegate that records each call and then runs [override]. */ + private fun overrideWith(override: (to: SubscriptionStatus) -> Unit) { + superwall.delegate = + object : SuperwallDelegate { + override fun subscriptionStatusDidChange( + from: SubscriptionStatus, + to: SubscriptionStatus, + ) { + calls += from.ids() to to.ids() + override(to) + } + } + } + + /** Entitlement ids of an Active status, an empty set for Inactive, null for Unknown. */ + private fun SubscriptionStatus.ids(): Set? = + when (this) { + is SubscriptionStatus.Active -> entitlements.map { it.id }.toSet() + is SubscriptionStatus.Inactive -> emptySet() + is SubscriptionStatus.Unknown -> null + } + + private fun awaitStatus(expected: Set) { + val deadline = System.currentTimeMillis() + 5_000 + while (System.currentTimeMillis() < deadline) { + if (superwall.subscriptionStatus.value.ids() == expected && calls.lastOrNull()?.second == expected) return + Thread.sleep(10) + } + assertEquals(expected, superwall.subscriptionStatus.value.ids()) + assertEquals("delegate was not told about the final status", expected, calls.lastOrNull()?.second) + } + + @Test + fun `delegate can add an entitlement to the status`() { + Given("a delegate that adds a custom entitlement whenever it is missing") { + overrideWith { to -> + if (to is SubscriptionStatus.Active && custom !in to.entitlements) { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(to.entitlements + custom)) + } + } + + When("the status becomes Active without it") { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(setOf(pro))) + awaitStatus(setOf("pro", "custom")) + + Then("the status and active entitlements include the added one") { + assertEquals(setOf("pro", "custom"), superwall.subscriptionStatus.value.ids()) + assertEquals(setOf("pro", "custom"), superwall.entitlements.active.map { it.id }.toSet()) + } + And("the delegate saw the original change, then its own override") { + assertEquals( + listOf(null to setOf("pro"), setOf("pro") to setOf("pro", "custom")), + calls.toList(), + ) + } + And("the overridden status is what gets persisted last") { + val writes = mutableListOf() + verify { storage.write(StoredSubscriptionStatus, capture(writes)) } + assertEquals(setOf("pro", "custom"), writes.last().ids()) + } + } + } + } + + @Test + fun `delegate can add an entitlement using the string overload`() { + Given("a delegate that re-sets the status by entitlement id") { + overrideWith { to -> + if (to is SubscriptionStatus.Active && custom !in to.entitlements) { + superwall.setSubscriptionStatus(*(to.entitlements.map { it.id } + "custom").toTypedArray()) + } + } + + When("the status becomes Active without it") { + superwall.setSubscriptionStatus("pro") + awaitStatus(setOf("pro", "custom")) + + Then("the status includes the added entitlement") { + assertEquals(setOf("pro", "custom"), superwall.subscriptionStatus.value.ids()) + } + } + } + } + + @Test + fun `delegate can remove an entitlement from the status`() { + Given("a delegate that strips the legacy entitlement") { + overrideWith { to -> + if (to is SubscriptionStatus.Active && legacy in to.entitlements) { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(to.entitlements - legacy)) + } + } + + When("the status becomes Active with it") { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(setOf(pro, legacy))) + awaitStatus(setOf("pro")) + + Then("the status no longer carries the removed entitlement") { + assertEquals(setOf("pro"), superwall.subscriptionStatus.value.ids()) + } + And("the delegate saw the original change, then its own override") { + assertEquals( + listOf(null to setOf("pro", "legacy"), setOf("pro", "legacy") to setOf("pro")), + calls.toList(), + ) + } + And("the narrowed status is what gets persisted last") { + val writes = mutableListOf() + verify { storage.write(StoredSubscriptionStatus, capture(writes)) } + assertEquals(setOf("pro"), writes.last().ids()) + } + // Active statuses only ever add to `entitlements.active`; it is cleared by + // Inactive/Unknown. Same as before the actor refactor. Pinned so a change + // here is a deliberate one. + And("entitlements.active still holds it until the status goes Inactive") { + assertEquals(setOf("pro", "legacy"), superwall.entitlements.active.map { it.id }.toSet()) + } + } + } + } + + @Test + fun `delegate can remove every entitlement by setting Inactive`() { + Given("a delegate that rejects any Active status") { + overrideWith { to -> + if (to is SubscriptionStatus.Active) { + superwall.setSubscriptionStatus(SubscriptionStatus.Inactive) + } + } + + When("the status becomes Active") { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(setOf(pro, legacy))) + awaitStatus(emptySet()) + + Then("the status is Inactive and nothing is active") { + assertTrue(superwall.subscriptionStatus.value is SubscriptionStatus.Inactive) + assertTrue(superwall.entitlements.active.isEmpty()) + } + And("the delegate saw the original change, then its own override") { + assertEquals( + listOf(null to setOf("pro", "legacy"), setOf("pro", "legacy") to emptySet()), + calls.toList(), + ) + } + } + } + } + + @Test + fun `delegate can replace one entitlement with another`() { + Given("a delegate that swaps legacy for custom") { + overrideWith { to -> + if (to is SubscriptionStatus.Active && legacy in to.entitlements) { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(to.entitlements - legacy + custom)) + } + } + + When("the status becomes Active with legacy") { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(setOf(pro, legacy))) + awaitStatus(setOf("pro", "custom")) + + Then("the status carries the replacement and not the original") { + assertEquals(setOf("pro", "custom"), superwall.subscriptionStatus.value.ids()) + } + } + } + } + + @Test + fun `delegate can grant entitlements when the status goes Inactive`() { + Given("a delegate that grants a custom entitlement to inactive users") { + overrideWith { to -> + if (to is SubscriptionStatus.Inactive) { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(setOf(custom))) + } + } + + When("the status becomes Inactive") { + superwall.setSubscriptionStatus(SubscriptionStatus.Inactive) + awaitStatus(setOf("custom")) + + Then("the status is Active with the granted entitlement") { + assertEquals(setOf("custom"), superwall.subscriptionStatus.value.ids()) + assertEquals(setOf("custom"), superwall.entitlements.active.map { it.id }.toSet()) + } + And("the delegate saw Inactive, then its own grant") { + assertEquals( + listOf(null to emptySet(), emptySet() to setOf("custom")), + calls.toList(), + ) + } + } + } + } + + @Test + fun `override is applied again on every later status change`() { + Given("a delegate that adds a custom entitlement whenever it is missing") { + overrideWith { to -> + if (to is SubscriptionStatus.Active && custom !in to.entitlements) { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(to.entitlements + custom)) + } + } + superwall.setSubscriptionStatus(SubscriptionStatus.Active(setOf(pro))) + awaitStatus(setOf("pro", "custom")) + + When("the SDK later sets a status without the custom entitlement") { + superwall.setSubscriptionStatus(SubscriptionStatus.Active(setOf(pro, legacy))) + awaitStatus(setOf("pro", "legacy", "custom")) + + Then("the override has been re-applied") { + assertEquals(setOf("pro", "legacy", "custom"), superwall.subscriptionStatus.value.ids()) + } + } + } + } +} diff --git a/superwall/src/test/java/com/superwall/sdk/config/ConfigManagerTest.kt b/superwall/src/test/java/com/superwall/sdk/config/ConfigManagerTest.kt index 5c1aa7efb..fa7d22357 100644 --- a/superwall/src/test/java/com/superwall/sdk/config/ConfigManagerTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/config/ConfigManagerTest.kt @@ -4,7 +4,6 @@ import android.content.Context import com.superwall.sdk.analytics.Tier import com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent import com.superwall.sdk.analytics.internal.trackable.TrackableSuperwallEvent -import com.superwall.sdk.config.models.ConfigState import com.superwall.sdk.config.options.SuperwallOptions import com.superwall.sdk.identity.IdentityManager import com.superwall.sdk.misc.Either @@ -29,7 +28,6 @@ import com.superwall.sdk.store.StoreManager import com.superwall.sdk.store.testmode.TestMode import com.superwall.sdk.store.testmode.TestModeBehavior import com.superwall.sdk.web.WebPaywallRedeemer -import com.superwall.sdk.models.assignment.Assignment import com.superwall.sdk.storage.DisableVerboseEvents import io.mockk.Runs import io.mockk.coEvery @@ -84,7 +82,6 @@ class ConfigManagerTest { val testMode: TestMode?, val tracked: CopyOnWriteArrayList, val statuses: MutableList, - val activateCalls: AtomicInteger, ) @Suppress("LongParameterList") @@ -164,7 +161,6 @@ class ConfigManagerTest { val tracked = CopyOnWriteArrayList() val statuses = mutableListOf() - val activateCalls = AtomicInteger(0) val options = SuperwallOptions().apply { @@ -191,9 +187,6 @@ class ConfigManagerTest { testMode = injectedTestMode, tracker = { tracked.add(it) }, setSubscriptionStatus = { statuses.add(it) }, - activateTestMode = { _, justActivated -> - if (justActivated) activateCalls.incrementAndGet() - }, identityManager = identityManager?.let { im -> { im } }, ) return Setup( @@ -208,7 +201,6 @@ class ConfigManagerTest { injectedTestMode, tracked, statuses, - activateCalls, ) } @@ -272,20 +264,55 @@ class ConfigManagerTest { fun `reevaluateTestMode activates when user now qualifies`() = runTest(timeout = 30.seconds) { val storageForTm = mockk(relaxed = true) - val testMode = TestMode(storage = storageForTm, isTestEnvironment = false) + val testMode = spyk(TestMode(storage = storageForTm, isTestEnvironment = false)) assertFalse(testMode.isTestMode) + setup( + backgroundScope, + testModeBehavior = TestModeBehavior.ALWAYS, + injectedTestMode = testMode, + ).manager.reevaluateTestMode(config = Config.stub(), appUserId = "anyone") + advanceUntilIdle() + + assertTrue(testMode.isTestMode) + coVerify(exactly = 1) { testMode.activate(any(), justActivated = true) } + } + + // ApplyConfig has its own flip-down branch (separate from ReevaluateTestMode): + // when a config arrives that no longer activates test mode, ApplyConfig must + // call clearTestModeState() and emit SubscriptionStatus.Inactive. + @Test + fun `ApplyConfig deactivates when prior testMode no longer qualifies under new config`() = + runTest(timeout = 30.seconds) { + val storageForTm = mockk(relaxed = true) + val testMode = spyk(TestMode(storage = storageForTm, isTestEnvironment = false)) + // Pre-activate so ApplyConfig sees wasTestMode=true. + testMode.evaluateTestMode( + Config.stub(), + "com.test", + null, + null, + testModeBehavior = TestModeBehavior.ALWAYS, + ) + assertTrue(testMode.isTestMode) + + // AUTOMATIC + Config.stub() (no userIds, matching bundleId) → deactivates. val s = setup( backgroundScope, - testModeBehavior = TestModeBehavior.ALWAYS, + testModeBehavior = TestModeBehavior.AUTOMATIC, injectedTestMode = testMode, ) - s.manager.reevaluateTestMode(config = Config.stub(), appUserId = "anyone") + s.manager.fetchConfiguration() advanceUntilIdle() - assertTrue(testMode.isTestMode) - assertEquals("activateTestMode lambda must fire once", 1, s.activateCalls.get()) + assertFalse("ApplyConfig must deactivate test mode", testMode.isTestMode) + verify(atLeast = 1) { testMode.clearTestModeState() } + assertTrue( + "Expected SubscriptionStatus.Inactive emitted from ApplyConfig flip-down", + s.statuses.any { it is SubscriptionStatus.Inactive }, + ) + coVerify(exactly = 0) { testMode.activate(any(), any()) } } @Test @@ -306,8 +333,8 @@ class ConfigManagerTest { assertFalse(testMode.isTestMode) verify(exactly = 0) { testMode.clearTestModeState() } + coVerify(exactly = 0) { testMode.activate(any(), any()) } assertTrue("No subscription status published on no-op", s.statuses.isEmpty()) - assertEquals("activateTestMode must not fire on no-op", 0, s.activateCalls.get()) } // Both reevaluateTestMode and ApplyConfig mutate TestMode.state. They @@ -1364,7 +1391,6 @@ class ConfigManagerTest { testMode = null, tracker = {}, setSubscriptionStatus = null, - activateTestMode = { _, _ -> }, ) mgr.fetchConfiguration() @@ -1410,7 +1436,6 @@ class ConfigManagerTest { testMode = null, tracker = {}, setSubscriptionStatus = null, - activateTestMode = { _, _ -> }, ) mgr.fetchConfiguration() @@ -1475,7 +1500,6 @@ class ConfigManagerTest { testMode = null, tracker = {}, setSubscriptionStatus = null, - activateTestMode = { _, _ -> }, awaitUtilNetwork = { awaitCalls.incrementAndGet() }, ) @@ -1684,7 +1708,6 @@ internal class ConfigManagerForTest( testMode: TestMode?, tracker: suspend (TrackableSuperwallEvent) -> Unit, setSubscriptionStatus: ((SubscriptionStatus) -> Unit)?, - activateTestMode: suspend (Config, Boolean) -> Unit, identityManager: (() -> IdentityManager)? = null, awaitUtilNetwork: suspend () -> Unit = {}, ) : ConfigManager( @@ -1706,6 +1729,5 @@ internal class ConfigManagerForTest( identityManager = identityManager, setSubscriptionStatus = setSubscriptionStatus, awaitUtilNetwork = awaitUtilNetwork, - activateTestMode = activateTestMode, actor = SequentialActor(ConfigState.None, CoroutineScope(Dispatchers.Unconfined)), ) diff --git a/superwall/src/test/java/com/superwall/sdk/config/ConfigStateReducerTest.kt b/superwall/src/test/java/com/superwall/sdk/config/ConfigStateReducerTest.kt index 8ca97c1a0..fbf282a7b 100644 --- a/superwall/src/test/java/com/superwall/sdk/config/ConfigStateReducerTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/config/ConfigStateReducerTest.kt @@ -1,6 +1,5 @@ package com.superwall.sdk.config -import com.superwall.sdk.config.models.ConfigState import com.superwall.sdk.models.config.Config import org.junit.Assert.assertEquals import org.junit.Assert.assertSame diff --git a/superwall/src/test/java/com/superwall/sdk/config/PaywallPreloadTest.kt b/superwall/src/test/java/com/superwall/sdk/config/PaywallPreloadTest.kt index fc932efd2..ab32885e2 100644 --- a/superwall/src/test/java/com/superwall/sdk/config/PaywallPreloadTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/config/PaywallPreloadTest.kt @@ -131,10 +131,10 @@ class PaywallPreloadTest { preload.removeUnusedPaywallVCsFromCache(oldConfig, newConfig) Then("only removed and changed, non-presented paywalls are cleared from cache") { - verify { paywallManager.removePaywallView("remove") } - verify { paywallManager.removePaywallView("changed") } - verify(exactly = 0) { paywallManager.removePaywallView("keep") } - verify(exactly = 0) { paywallManager.removePaywallView("presented") } + coVerify { paywallManager.removePaywallView("remove") } + coVerify { paywallManager.removePaywallView("changed") } + coVerify(exactly = 0) { paywallManager.removePaywallView("keep") } + coVerify(exactly = 0) { paywallManager.removePaywallView("presented") } } } } diff --git a/superwall/src/test/java/com/superwall/sdk/misc/AwaitFirstValidConfigTest.kt b/superwall/src/test/java/com/superwall/sdk/misc/AwaitFirstValidConfigTest.kt index fc0c396ca..060aadc13 100644 --- a/superwall/src/test/java/com/superwall/sdk/misc/AwaitFirstValidConfigTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/misc/AwaitFirstValidConfigTest.kt @@ -1,6 +1,6 @@ package com.superwall.sdk.misc -import com.superwall.sdk.config.models.ConfigState +import com.superwall.sdk.config.ConfigState import com.superwall.sdk.models.config.Config import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.flow.flowOf diff --git a/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallManagerExperimentIsolationTest.kt b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallManagerExperimentIsolationTest.kt index f690dec4a..cf86130d3 100644 --- a/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallManagerExperimentIsolationTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallManagerExperimentIsolationTest.kt @@ -92,7 +92,7 @@ class PaywallManagerExperimentIsolationTest { val cache = mockk(relaxed = true) { every { getPaywallView(any()) } answers { cachedView } - every { save(any(), any()) } answers { cachedView = firstArg() } + coEvery { save(any(), any()) } answers { cachedView = firstArg() } } val deviceInfo = mockk { every { locale } returns "en_US" } val managerFactory = diff --git a/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallManagerTest.kt b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallManagerTest.kt index 2c61dba8d..7a4f4585e 100644 --- a/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallManagerTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallManagerTest.kt @@ -13,6 +13,7 @@ import com.superwall.sdk.paywall.view.delegate.PaywallLoadingState import com.superwall.sdk.paywall.view.delegate.PaywallViewDelegateAdapter import io.mockk.Runs import io.mockk.coEvery +import io.mockk.coVerify import io.mockk.every import io.mockk.just import io.mockk.mockk @@ -70,14 +71,15 @@ class PaywallManagerTest { } @Test - fun test_removePaywallView_callsCacheRemove() { - val identifier: PaywallIdentifier = "test_paywall" - every { cache.removePaywallView(any()) } just Runs + fun test_removePaywallView_callsCacheRemove() = + runTest { + val identifier: PaywallIdentifier = "test_paywall" + coEvery { cache.removePaywallView(any()) } just Runs - paywallManager.removePaywallView(identifier) + paywallManager.removePaywallView(identifier) - verify { cache.removePaywallView(identifier) } - } + coVerify { cache.removePaywallView(identifier) } + } @Test fun test_resetCache_destroysWebviewsAndClearsCache() = @@ -89,13 +91,13 @@ class PaywallManagerTest { every { mockView2.destroyWebview() } just Runs every { cache.getAllPaywallViews() } returns listOf(mockView1, mockView2) every { cache.activePaywallVcKey } returns null - every { cache.removeAll() } just Runs + coEvery { cache.removeAll() } just Runs paywallManager.resetCache() verify { mockView1.destroyWebview() } verify { mockView2.destroyWebview() } - verify { cache.removeAll() } + coVerify { cache.removeAll() } } @Test @@ -119,13 +121,13 @@ class PaywallManagerTest { coEvery { paywallRequestManager.getPaywall(any(), any()) } returns Either.Success(paywall) every { cache.getPaywallView(any()) } returns null coEvery { factory.makePaywallView(any(), any(), any()) } returns mockView - every { cache.save(any(), any()) } just Runs + coEvery { cache.save(any(), any()) } just Runs val result = paywallManager.getPaywallView(request, true, false, null) assertTrue(result is Either.Success) assertEquals(mockView, (result as Either.Success).value) - verify { cache.save(mockView, "test_paywall") } + coVerify { cache.save(mockView, "test_paywall") } } @Test @@ -260,7 +262,7 @@ class PaywallManagerTest { coEvery { paywallRequestManager.getPaywall(any(), any()) } returns Either.Success(paywall) coEvery { factory.makePaywallView(any(), any(), any()) } returns mockView - every { cache.save(any(), any()) } just Runs + coEvery { cache.save(any(), any()) } just Runs val result = paywallManager.getPaywallView(request, true, false, null) @@ -288,7 +290,7 @@ class PaywallManagerTest { coEvery { paywallRequestManager.getPaywall(any(), any()) } returns Either.Success(paywall) every { cache.getPaywallView(any()) } returns null coEvery { factory.makePaywallView(any(), any(), any()) } returns mockView - every { cache.save(any(), any()) } just Runs + coEvery { cache.save(any(), any()) } just Runs paywallManager.getPaywallView(request, isForPresentation = true, isPreloading = false, null) @@ -314,7 +316,7 @@ class PaywallManagerTest { coEvery { paywallRequestManager.getPaywall(any(), any()) } returns Either.Success(paywall) every { cache.getPaywallView(any()) } returns null coEvery { factory.makePaywallView(any(), any(), any()) } returns mockView - every { cache.save(any(), any()) } just Runs + coEvery { cache.save(any(), any()) } just Runs paywallManager.getPaywallView(request, isForPresentation = false, isPreloading = false, null) diff --git a/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt new file mode 100644 index 000000000..22e499404 --- /dev/null +++ b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt @@ -0,0 +1,587 @@ +package com.superwall.sdk.paywall.manager + +import android.view.View +import com.superwall.sdk.Given +import com.superwall.sdk.Then +import com.superwall.sdk.When +import com.superwall.sdk.misc.ActivityProvider +import com.superwall.sdk.misc.primitives.SequentialActor +import com.superwall.sdk.network.device.DeviceHelper +import com.superwall.sdk.paywall.view.LoadingView +import com.superwall.sdk.paywall.view.PaywallView +import com.superwall.sdk.paywall.view.ShimmerView +import com.superwall.sdk.paywall.view.ViewStorage +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.SupervisorJob +import kotlinx.coroutines.async +import kotlinx.coroutines.awaitAll +import kotlinx.coroutines.cancel +import kotlinx.coroutines.launch +import kotlinx.coroutines.test.runTest +import org.junit.After +import org.junit.Assert.assertEquals +import org.junit.Assert.assertNotNull +import org.junit.Assert.assertNull +import org.junit.Assert.assertSame +import org.junit.Assert.assertTrue +import org.junit.Before +import org.junit.Test +import org.junit.runner.RunWith +import org.robolectric.RobolectricTestRunner +import org.robolectric.RuntimeEnvironment +import org.robolectric.annotation.Config +import java.util.concurrent.ConcurrentHashMap + +@RunWith(RobolectricTestRunner::class) +@Config(sdk = [33]) +class PaywallViewCacheTest { + private lateinit var appCtx: android.content.Context + private lateinit var activityProvider: ActivityProvider + private lateinit var deviceHelper: DeviceHelper + private lateinit var storage: ViewStorage + + private fun keyOf(id: String) = PaywallCacheLogic.key(id, "en_US") + + private lateinit var actorScope: CoroutineScope + + private fun newCache(): PaywallViewCache = + PaywallViewCache( + appCtx, + storage, + activityProvider, + deviceHelper, + actor = SequentialActor(PaywallCacheState(), actorScope), + ) + + @Before + fun setup() { + appCtx = RuntimeEnvironment.getApplication() + activityProvider = + mockk { + every { getCurrentActivity() } returns null + } + deviceHelper = + mockk { + every { locale } returns "en_US" + } + storage = + object : ViewStorage { + override val views = ConcurrentHashMap() + } + actorScope = CoroutineScope(SupervisorJob() + Dispatchers.IO) + } + + @After + fun tearDown() { + actorScope.cancel() + } + + // ------------------------------------------------------------------- + // Init / pre-population + // ------------------------------------------------------------------- + + @Test + fun `acquireLoadingView creates and stores it under LoadingView TAG`() { + Given("a fresh cache") { + val cache = newCache() + + When("acquireLoadingView is called") { + val view = cache.acquireLoadingView() + + Then("the view exists in storage under its tag") { + assertNotNull(view) + assertNotNull(storage.retrieveView(LoadingView.TAG)) + } + } + } + } + + // ------------------------------------------------------------------- + // save / get + // ------------------------------------------------------------------- + + @Test + fun `save then getPaywallView returns the view synchronously`() = + runTest { + Given("a cache and a paywall view") { + val cache = newCache() + val view = mockk(relaxed = true) + + When("saved and immediately fetched") { + cache.save(view, "paywall_a") + val result = cache.getPaywallView(keyOf("paywall_a")) + + Then("the same instance is returned without delay") { + assertSame(view, result) + } + } + } + } + + @Test + fun `getPaywallView returns null for unknown key`() { + Given("a cache with no saved paywalls") { + val cache = newCache() + + Then("looking up a missing key returns null") { + assertNull(cache.getPaywallView("missing")) + } + } + } + + @Test + fun `save uses locale-aware cache key`() = + runTest { + Given("a device locale of en_US") { + val cache = newCache() + val view = mockk(relaxed = true) + + When("save is called with identifier 'foo'") { + cache.save(view, "foo") + + Then("the view is stored under 'foo_en_US'") { + assertSame(view, cache.getPaywallView("foo_en_US")) + } + } + } + } + + @Test + fun `saving same identifier twice keeps the latest view`() = + runTest { + Given("two views saved under the same identifier") { + val cache = newCache() + val first = mockk(relaxed = true) + val second = mockk(relaxed = true) + + cache.save(first, "dup") + cache.save(second, "dup") + + Then("getPaywallView returns the second") { + assertSame(second, cache.getPaywallView(keyOf("dup"))) + } + } + } + + // ------------------------------------------------------------------- + // activePaywallVcKey / activePaywallView + // ------------------------------------------------------------------- + + @Test + fun `activePaywallVcKey defaults to null`() { + val cache = newCache() + assertNull(cache.activePaywallVcKey) + assertNull(cache.activePaywallView) + } + + @Test + fun `setting activePaywallVcKey is observable on subsequent reads`() { + Given("a cache") { + val cache = newCache() + + When("the active key is set") { + cache.activePaywallVcKey = "abc" + + Then("the read returns the same value") { + assertEquals("abc", cache.activePaywallVcKey) + } + } + } + } + + @Test + fun `activePaywallView returns the view stored under activePaywallVcKey`() = + runTest { + Given("a saved paywall and matching active key") { + val cache = newCache() + val view = mockk(relaxed = true) + cache.save(view, "foo") + cache.activePaywallVcKey = keyOf("foo") + + Then("activePaywallView returns it") { + assertSame(view, cache.activePaywallView) + } + } + } + + @Test + fun `activePaywallView is null when activeKey points at non-PaywallView`() { + Given("activeKey set to LoadingView's tag") { + val cache = newCache() + cache.activePaywallVcKey = LoadingView.TAG + + Then("activePaywallView is null (cast guarded)") { + assertNull(cache.activePaywallView) + } + } + } + + @Test + fun `activePaywallView is null when key has no entry`() { + val cache = newCache() + cache.activePaywallVcKey = "ghost" + assertNull(cache.activePaywallView) + } + + // ------------------------------------------------------------------- + // getAllPaywallViews / entries + // ------------------------------------------------------------------- + + @Test + fun `getAllPaywallViews excludes loading and shimmer views`() = + runTest { + Given("two saved paywalls") { + val cache = newCache() + val a = mockk(relaxed = true) + val b = mockk(relaxed = true) + cache.save(a, "a") + cache.save(b, "b") + + Then("only the paywall views are returned") { + val views = cache.getAllPaywallViews() + assertEquals(2, views.size) + assertTrue(views.contains(a)) + assertTrue(views.contains(b)) + } + } + } + + @Test + fun `getAllPaywallViews is empty when no paywalls saved`() { + val cache = newCache() + assertTrue(cache.getAllPaywallViews().isEmpty()) + } + + // ------------------------------------------------------------------- + // removePaywallView / removeAll + // ------------------------------------------------------------------- + + @Test + fun `removePaywallView removes only that identifier`() = + runTest { + Given("two saved paywalls") { + val cache = newCache() + val a = mockk(relaxed = true) + val b = mockk(relaxed = true) + cache.save(a, "a") + cache.save(b, "b") + + When("one is removed") { + cache.removePaywallView("a") + + Then("only the other remains") { + assertNull(cache.getPaywallView(keyOf("a"))) + assertSame(b, cache.getPaywallView(keyOf("b"))) + } + } + } + } + + @Test + fun `removeAll preserves the active key entry`() = + runTest { + Given("multiple saved paywalls and an active key") { + val cache = newCache() + val a = mockk(relaxed = true) + val b = mockk(relaxed = true) + cache.save(a, "a") + cache.save(b, "b") + cache.activePaywallVcKey = keyOf("a") + + When("removeAll is called") { + cache.removeAll() + + Then("the active entry survives") { + assertSame(a, cache.getPaywallView(keyOf("a"))) + } + Then("the inactive entry is gone") { + assertNull(cache.getPaywallView(keyOf("b"))) + } + } + } + } + + @Test + fun `removeAll with no active key clears every entry`() = + runTest { + Given("two saved paywalls and no active key") { + val cache = newCache() + cache.save(mockk(relaxed = true), "a") + cache.save(mockk(relaxed = true), "b") + + When("removeAll is called") { + cache.removeAll() + + Then("no paywall views remain") { + assertTrue(cache.getAllPaywallViews().isEmpty()) + } + } + } + } + + // ------------------------------------------------------------------- + // acquireLoadingView / acquireShimmerView + // ------------------------------------------------------------------- + + @Test + fun `acquireLoadingView returns the cached instance on repeat calls`() { + val cache = newCache() + val first = cache.acquireLoadingView() + val second = cache.acquireLoadingView() + assertSame(first, second) + } + + @Test + fun `acquireLoadingView is atomic across concurrent callers`() = + runTest { + Given("many concurrent acquireLoadingView calls on a fresh cache") { + val cache = newCache() + val results = + (0 until 32) + .map { async(Dispatchers.Default) { cache.acquireLoadingView() } } + .awaitAll() + + Then("every caller receives the same canonical instance") { + val canonical = results.first() + assertTrue(results.all { it === canonical }) + } + Then("the canonical instance is what's stored") { + assertSame(results.first() as View, storage.retrieveView(LoadingView.TAG)) + } + } + } + + @Test + fun `cache hydrates from existing viewStorage entries on construction`() { + Given("a viewStorage already populated before cache construction") { + val pre = mockk(relaxed = true) + storage.storeView(keyOf("pre"), pre) + + When("a fresh cache is built") { + val cache = newCache() + + Then("the existing entry is visible to the cache") { + assertSame(pre, cache.getPaywallView(keyOf("pre"))) + } + } + } + } + + // ------------------------------------------------------------------- + // Concurrency / ordering + // ------------------------------------------------------------------- + + @Test + fun `concurrent saves from many coroutines all land`() = + runTest { + Given("100 saves from background coroutines") { + val cache = newCache() + val views = (0 until 100).map { mockk(relaxed = true) } + + val jobs = + views.mapIndexed { i, v -> + launch(Dispatchers.Default) { cache.save(v, "p_$i") } + } + jobs.forEach { it.join() } + + Then("all entries are retrievable") { + views.forEachIndexed { i, v -> + assertSame("missing $i", v, cache.getPaywallView(keyOf("p_$i"))) + } + } + } + } + + @Test + fun `concurrent activeKey writes leave a consistent final value`() = + runTest { + Given("repeated concurrent activeKey assignments") { + val cache = newCache() + + val jobs = + (0 until 50).map { i -> + launch(Dispatchers.Default) { cache.activePaywallVcKey = "k_$i" } + } + jobs.forEach { it.join() } + + Then("the final read returns one of the assigned values") { + val final = cache.activePaywallVcKey + assertTrue(final in (0 until 50).map { "k_$it" }) + } + } + } + + @Test + fun `interleaved saves and removes converge to a stable state`() = + runTest { + Given("concurrent saves and removes on the same identifiers") { + val cache = newCache() + val ids = (0 until 20).map { "id_$it" } + + val savers = + ids.map { id -> + async(Dispatchers.Default) { cache.save(mockk(relaxed = true), id) } + } + val removers = + ids.map { id -> + async(Dispatchers.Default) { cache.removePaywallView(id) } + } + (savers + removers).awaitAll() + + Then("cache state and viewStorage agree on every key") { + // Which of save/remove wins per id depends on interleaving, but the + // two stores must end up agreeing on it. + ids.forEach { id -> + val key = keyOf(id) + assertSame(storage.retrieveView(key), cache.getPaywallView(key)) + } + assertEquals( + storage.all().filterIsInstance().size, + cache.getAllPaywallViews().size, + ) + } + } + } + + @Test + fun `removeAll then save then read returns the new view`() = + runTest { + Given("a populated cache that has been cleared") { + val cache = newCache() + cache.save(mockk(relaxed = true), "old") + cache.removeAll() + + When("a new view is saved after the clear") { + val fresh = mockk(relaxed = true) + cache.save(fresh, "new") + + Then("the new view is readable") { + assertSame(fresh, cache.getPaywallView(keyOf("new"))) + } + } + } + } + + // ------------------------------------------------------------------- + // External writers (SuperwallPaywallActivity, DebugView) + // ------------------------------------------------------------------- + + @Test + fun `removeView evicts a saved paywall from both cache and viewStorage`() = + runTest { + Given("a saved paywall whose activity launch then fails") { + val cache = newCache() + val view = mockk(relaxed = true) + cache.save(view, "p1") + + When("the launch-failure path removes its key") { + cache.removeView(keyOf("p1")) + + Then("the next lookup misses, forcing a fresh view") { + assertNull(cache.getPaywallView(keyOf("p1"))) + assertNull(storage.retrieveView(keyOf("p1"))) + } + } + } + } + + @Test + fun `storeView is visible to cache reads and viewStorage immediately`() { + Given("a view stored under an activity key") { + val cache = newCache() + val view = mockk(relaxed = true) + + When("storeView is called") { + cache.storeView("activity-key", view) + + Then("both the cache and viewStorage return it synchronously") { + assertSame(view, cache.getPaywallView("activity-key")) + assertSame(view, storage.retrieveView("activity-key")) + } + } + } + } + + @Test + fun `removeAll sweeps views stored through storeView`() = + runTest { + Given("a debug view stored under an arbitrary key") { + val cache = newCache() + cache.storeView("debug-key", View(appCtx)) + + When("removeAll runs") { + cache.removeAll() + + Then("the debug view is gone from both stores") { + assertNull(cache.entries["debug-key"]) + assertNull(storage.retrieveView("debug-key")) + } + } + } + } + + // ------------------------------------------------------------------- + // Loading/shimmer availability for startWithView without present() + // ------------------------------------------------------------------- + + @Test + fun `loading and shimmer are not in viewStorage until acquired`() { + Given("a cold cache, as seen by getPaywall() + startWithView()") { + newCache() + + Then("readers must acquire through the cache instead of viewStorage") { + assertNull(storage.retrieveView(LoadingView.TAG)) + assertNull(storage.retrieveView(ShimmerView.TAG)) + } + } + } + + @Test + fun `acquire recreates loading and shimmer after removeAll evicts them`() = + runTest { + Given("acquired loading and shimmer views") { + val cache = newCache() + val loading = cache.acquireLoadingView() + val shimmer = cache.acquireShimmerView() + + When("removeAll evicts them and they are acquired again") { + cache.removeAll() + val newLoading = cache.acquireLoadingView() + val newShimmer = cache.acquireShimmerView() + + Then("fresh instances are stored under their tags") { + assertTrue(newLoading !== loading) + assertTrue(newShimmer !== shimmer) + assertSame(newLoading, storage.retrieveView(LoadingView.TAG)) + assertSame(newShimmer, storage.retrieveView(ShimmerView.TAG)) + } + } + } + } + + @Test + fun `registry writes reach both the cache and viewStorage`() { + Given("the registry view of a cache") { + val cache = newCache() + val registry = cache.asRegistry() + val view = mockk(relaxed = true) + + When("a view is stored and then removed through the registry") { + registry.storeView("activity-key", view) + val stored = registry.retrieveView("activity-key") + val inCache = cache.getPaywallView("activity-key") + registry.removeView("activity-key") + + Then("both stores saw each write") { + assertSame(view, stored) + assertSame(view, inCache) + assertNull(cache.getPaywallView("activity-key")) + assertNull(storage.retrieveView("activity-key")) + } + } + } + } +} diff --git a/superwall/src/test/java/com/superwall/sdk/paywall/presentation/internal/operators/WaitForSubsStatusAndConfigTest.kt b/superwall/src/test/java/com/superwall/sdk/paywall/presentation/internal/operators/WaitForSubsStatusAndConfigTest.kt index e985671f8..2a325dfd3 100644 --- a/superwall/src/test/java/com/superwall/sdk/paywall/presentation/internal/operators/WaitForSubsStatusAndConfigTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/paywall/presentation/internal/operators/WaitForSubsStatusAndConfigTest.kt @@ -4,7 +4,7 @@ import com.superwall.sdk.Given import com.superwall.sdk.Then import com.superwall.sdk.When import com.superwall.sdk.analytics.internal.track -import com.superwall.sdk.config.models.ConfigState +import com.superwall.sdk.config.ConfigState import com.superwall.sdk.dependencies.DependencyContainer import com.superwall.sdk.models.config.Config import com.superwall.sdk.models.entitlements.SubscriptionStatus diff --git a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt new file mode 100644 index 000000000..2663ab8ec --- /dev/null +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt @@ -0,0 +1,1330 @@ +package com.superwall.sdk.store + +import com.superwall.sdk.And +import com.superwall.sdk.Given +import com.superwall.sdk.Then +import com.superwall.sdk.When +import com.superwall.sdk.models.customer.CustomerInfo +import com.superwall.sdk.models.entitlements.Entitlement +import com.superwall.sdk.models.entitlements.SubscriptionStatus +import com.superwall.sdk.models.internal.WebRedemptionResponse +import com.superwall.sdk.models.product.Store +import com.superwall.sdk.storage.LatestRedemptionResponse +import com.superwall.sdk.storage.Storage +import com.superwall.sdk.storage.StoredEntitlementsByProductId +import com.superwall.sdk.storage.StoredSubscriptionStatus +import com.superwall.sdk.store.abstractions.product.receipt.LatestSubscriptionState +import io.mockk.every +import io.mockk.mockk +import io.mockk.verify +import kotlinx.coroutines.launch +import kotlinx.coroutines.test.UnconfinedTestDispatcher +import kotlinx.coroutines.test.advanceUntilIdle +import kotlinx.coroutines.test.runTest +import org.junit.Assert.assertEquals +import org.junit.Assert.assertFalse +import org.junit.Assert.assertTrue +import org.junit.Test +import java.util.Date + +/** + * Comprehensive tests for the Entitlements class external API. + * These tests are designed to guarantee correctness after refactoring. + * + * Covers: + * - Initialization (clean, cached, corrupted) + * - setSubscriptionStatus (all transitions, edge cases) + * - Property computations (active, inactive, all, web) + * - Product ID lookup (exact, partial, fallback chains) + * - addEntitlementsByProductId + * - entitlementsByProductId snapshot + * - byProductIds (batch) + * - activeDeviceEntitlements lifecycle + * - Status flow persistence + * - Multi-step state transitions + * - Deduplication and merge priority + */ +class EntitlementsRefactorSafetyTest { + private fun mockStorage( + storedStatus: SubscriptionStatus? = null, + storedProductEntitlements: Map>? = null, + redemptionResponse: WebRedemptionResponse? = null, + ): Storage = + mockk(relaxUnitFun = true) { + every { read(StoredSubscriptionStatus) } returns storedStatus + every { read(StoredEntitlementsByProductId) } returns storedProductEntitlements + every { read(LatestRedemptionResponse) } returns redemptionResponse + } + + private fun webRedemption(vararg entitlements: Entitlement): WebRedemptionResponse = + WebRedemptionResponse( + codes = emptyList(), + allCodes = emptyList(), + customerInfo = + CustomerInfo( + subscriptions = emptyList(), + nonSubscriptions = emptyList(), + userId = "testUser", + entitlements = entitlements.toList(), + isPlaceholder = false, + ), + ) + + // ========================================== + // Initialization Edge Cases + // ========================================== + + @Test + fun `init with no stored data starts with Unknown status and empty collections`() = + runTest { + Given("storage has no cached data") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + + Then("status should be Unknown") { + assertTrue(entitlements.status.value is SubscriptionStatus.Unknown) + } + And("all collections should be empty") { + assertTrue(entitlements.active.isEmpty()) + assertTrue(entitlements.inactive.isEmpty()) + assertTrue(entitlements.all.isEmpty()) + assertTrue(entitlements.web.isEmpty()) + assertTrue(entitlements.entitlementsByProductId.isEmpty()) + assertTrue(entitlements.activeDeviceEntitlements.isEmpty()) + } + } + } + + @Test + fun `init with corrupted StoredSubscriptionStatus resets to Unknown`() = + runTest { + Given("storage throws ClassCastException for StoredSubscriptionStatus") { + val storage = + mockk(relaxUnitFun = true) { + every { read(StoredSubscriptionStatus) } throws ClassCastException("corrupted") + every { read(StoredEntitlementsByProductId) } returns null + every { read(LatestRedemptionResponse) } returns null + } + + When("Entitlements is initialized") { + val entitlements = + makeEntitlements(storage, backgroundScope) + + Then("corrupted status should be deleted from storage") { + verify { storage.delete(StoredSubscriptionStatus) } + } + And("status should be set to Unknown") { + assertTrue(entitlements.status.value is SubscriptionStatus.Unknown) + } + } + } + } + + @Test + fun `init keeps status-only entitlements in all when product entitlements are stored`() = + runTest { + Given("a stored Active status and stored product entitlements that do not overlap") { + val statusOnly = Entitlement("status_only") + val productOnly = Entitlement("product_only") + val storage = + mockStorage( + storedStatus = SubscriptionStatus.Active(setOf(statusOnly)), + storedProductEntitlements = mapOf("product_1" to setOf(productOnly)), + ) + + When("Entitlements is created on a cold start") { + val entitlements = makeEntitlements(storage, backgroundScope) + + Then("all contains both the status and the product entitlements") { + assertEquals( + setOf("status_only", "product_only"), + entitlements.all.map { it.id }.toSet(), + ) + } + And("every active entitlement is also in all") { + val allIds = entitlements.all.map { it.id }.toSet() + assertTrue(entitlements.active.all { it.id in allIds }) + } + } + } + } + + @Test + fun `addEntitlementsByProductId keeps status-only entitlements in all`() = + runTest { + Given("an Active status with an entitlement not tied to any product") { + val statusOnly = Entitlement("status_only") + val productOnly = Entitlement("product_only") + val storage = mockStorage(storedStatus = SubscriptionStatus.Active(setOf(statusOnly))) + val entitlements = makeEntitlements(storage, backgroundScope) + + When("product entitlements are added afterwards") { + entitlements.addEntitlementsByProductId(mapOf("product_1" to setOf(productOnly))) + + Then("all contains both the status and the product entitlements") { + assertEquals( + setOf("status_only", "product_only"), + entitlements.all.map { it.id }.toSet(), + ) + } + And("every active entitlement is also in all") { + val allIds = entitlements.all.map { it.id }.toSet() + assertTrue(entitlements.active.all { it.id in allIds }) + } + } + } + } + + @Test + fun `init with corrupted StoredEntitlementsByProductId deletes and continues`() = + runTest { + Given("storage throws ClassCastException for StoredEntitlementsByProductId") { + val storage = + mockk(relaxUnitFun = true) { + every { read(StoredSubscriptionStatus) } returns null + every { read(StoredEntitlementsByProductId) } throws ClassCastException("corrupted") + every { read(LatestRedemptionResponse) } returns null + } + + When("Entitlements is initialized") { + val entitlements = + makeEntitlements(storage, backgroundScope) + + Then("corrupted entitlements-by-product should be deleted") { + verify { storage.delete(StoredEntitlementsByProductId) } + } + And("entitlementsByProductId should be empty") { + assertTrue(entitlements.entitlementsByProductId.isEmpty()) + } + } + } + } + + @Test + fun `init with stored Inactive status restores Inactive`() = + runTest { + Given("storage contains Inactive status") { + val storage = mockStorage(storedStatus = SubscriptionStatus.Inactive) + val entitlements = + makeEntitlements(storage, backgroundScope) + + Then("status should be Inactive") { + assertTrue(entitlements.status.value is SubscriptionStatus.Inactive) + } + And("active and inactive should be empty") { + assertTrue(entitlements.active.isEmpty()) + assertTrue(entitlements.inactive.isEmpty()) + } + } + } + + @Test + fun `init with stored product entitlements restores them`() = + runTest { + Given("storage contains product entitlements") { + val e1 = Entitlement("premium") + val productMap = mapOf("prod1" to setOf(e1)) + val storage = mockStorage(storedProductEntitlements = productMap) + + When("Entitlements is initialized") { + val entitlements = + makeEntitlements(storage, backgroundScope) + + Then("entitlementsByProductId should contain the stored mappings") { + assertEquals(productMap, entitlements.entitlementsByProductId) + } + And("all should include entitlements from product map") { + assertTrue(entitlements.all.contains(e1)) + } + } + } + } + + // ========================================== + // setSubscriptionStatus - Active Entitlement Filtering + // ========================================== + + @Test + fun `setSubscriptionStatus Active only adds isActive entitlements to backingActive`() = + runTest { + Given("a mix of active and inactive entitlements") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val activeE = Entitlement("active_one", isActive = true) + val inactiveE = Entitlement("inactive_one", isActive = false) + + When("setting Active status with both") { + entitlements.setSubscriptionStatus( + SubscriptionStatus.Active(setOf(activeE, inactiveE)), + ) + + Then("active should only contain the isActive entitlement") { + assertTrue(entitlements.active.any { it.id == "active_one" }) + } + And("the inactive entitlement should not be in active") { + assertFalse(entitlements.active.any { it.id == "inactive_one" && !it.isActive }) + } + And("all should contain both") { + assertTrue(entitlements.all.any { it.id == "active_one" }) + assertTrue(entitlements.all.any { it.id == "inactive_one" }) + } + } + } + } + + @Test + fun `setSubscriptionStatus Active with only inactive entitlements stays Active`() = + runTest { + Given("entitlements that are all inactive") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val inactiveE = + Entitlement( + id = "expired", + type = Entitlement.Type.SERVICE_LEVEL, + isActive = false, + ) + + When("setting Active status with only inactive entitlements") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(inactiveE))) + + Then("status should remain Active since set is not empty") { + // The code only checks entitlements.isEmpty(), not isActive + assertTrue(entitlements.status.value is SubscriptionStatus.Active) + } + } + } + } + + // ========================================== + // setSubscriptionStatus - State Transitions + // ========================================== + + @Test + fun `transition Active to Active replaces entitlements additively`() = + runTest { + Given("entitlements with Active status") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val e1 = Entitlement("first") + val e2 = Entitlement("second") + + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e1))) + + When("setting Active with different entitlements") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e2))) + + Then("active should contain both since backingActive uses addAll") { + assertTrue(entitlements.active.any { it.id == "first" }) + assertTrue(entitlements.active.any { it.id == "second" }) + } + And("status value should reflect latest set") { + val status = entitlements.status.value as SubscriptionStatus.Active + assertTrue(status.entitlements.any { it.id == "second" }) + } + } + } + } + + @Test + fun `transition Active to Inactive to Active restores correctly`() = + runTest { + Given("entitlements cycling through states") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val e1 = Entitlement("premium") + + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e1))) + + When("going Inactive then Active again") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) + + Then("after Inactive, active should be empty") { + assertTrue(entitlements.active.isEmpty()) + } + + val e2 = Entitlement("gold") + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e2))) + + And("after re-activation, only new entitlements should be active") { + assertTrue(entitlements.active.any { it.id == "gold" }) + // e1 was cleared by Inactive + assertFalse(entitlements.active.any { it.id == "premium" }) + } + } + } + } + + @Test + fun `transition Active to Unknown clears everything`() = + runTest { + Given("entitlements in Active state with device entitlements") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("a")))) + entitlements.activeDeviceEntitlements = setOf(Entitlement("device")) + + When("setting Unknown") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Unknown) + + Then("backingActive and activeDeviceEntitlements should be cleared") { + assertTrue(entitlements.active.isEmpty()) + assertTrue(entitlements.activeDeviceEntitlements.isEmpty()) + } + And("status should be Unknown") { + assertTrue(entitlements.status.value is SubscriptionStatus.Unknown) + } + } + } + } + + @Test + fun `transition Unknown to Inactive keeps collections empty`() = + runTest { + Given("entitlements in Unknown state") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("setting Inactive from Unknown") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) + + Then("all collections should remain empty") { + assertTrue(entitlements.active.isEmpty()) + assertTrue(entitlements.inactive.isEmpty()) + assertTrue(entitlements.status.value is SubscriptionStatus.Inactive) + } + } + } + } + + @Test + fun `multiple rapid state transitions end in correct final state`() = + runTest { + Given("entitlements subjected to rapid transitions") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val e1 = Entitlement("a") + val e2 = Entitlement("b") + val e3 = Entitlement("c") + + When("cycling through Active, Inactive, Unknown, Active") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e1))) + entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) + entitlements.setSubscriptionStatus(SubscriptionStatus.Unknown) + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e3))) + + Then("final status should be Active with e3") { + val status = entitlements.status.value + assertTrue(status is SubscriptionStatus.Active) + assertTrue(entitlements.active.any { it.id == "c" }) + } + And("e1 and e2 should not be in active (cleared by Inactive/Unknown)") { + assertFalse(entitlements.active.any { it.id == "a" }) + assertFalse(entitlements.active.any { it.id == "b" }) + } + } + } + } + + // ========================================== + // activeDeviceEntitlements Lifecycle + // ========================================== + + @Test + fun `activeDeviceEntitlements cleared on Unknown status`() = + runTest { + Given("entitlements with active device entitlements") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + entitlements.activeDeviceEntitlements = setOf(Entitlement("device_premium")) + + When("setting Unknown status") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Unknown) + + Then("activeDeviceEntitlements should be cleared") { + assertTrue(entitlements.activeDeviceEntitlements.isEmpty()) + } + } + } + } + + @Test + fun `activeDeviceEntitlements setter replaces not appends`() = + runTest { + Given("entitlements with existing device entitlements") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + entitlements.activeDeviceEntitlements = setOf(Entitlement("old")) + + When("setting new device entitlements") { + entitlements.activeDeviceEntitlements = setOf(Entitlement("new")) + + Then("only the new entitlement should be present") { + assertEquals(1, entitlements.activeDeviceEntitlements.size) + assertTrue(entitlements.activeDeviceEntitlements.any { it.id == "new" }) + assertFalse(entitlements.activeDeviceEntitlements.any { it.id == "old" }) + } + } + } + } + + @Test + fun `activeDeviceEntitlements do not persist to backingActive on Inactive`() = + runTest { + Given("device entitlements set, then status goes Inactive") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + entitlements.activeDeviceEntitlements = setOf(Entitlement("device")) + + When("setting Inactive") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) + + Then("active should be empty (device entitlements cleared by Inactive)") { + assertTrue(entitlements.active.isEmpty()) + } + } + } + } + + // ========================================== + // Property Computations - active, inactive, all + // ========================================== + + @Test + fun `all property combines _all, entitlementsByProduct values, and web`() = + runTest { + Given("entitlements from all three backing sources") { + val webE = Entitlement("web", isActive = true, store = Store.STRIPE) + val storage = + mockStorage( + storedProductEntitlements = mapOf("prod1" to setOf(Entitlement("from_product"))), + redemptionResponse = webRedemption(webE), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("from_status")))) + + When("accessing all property") { + val all = entitlements.all + + Then("it should contain entitlements from all three sources") { + assertTrue(all.any { it.id == "from_status" }) + assertTrue(all.any { it.id == "from_product" }) + assertTrue(all.any { it.id == "web" }) + } + } + } + } + + @Test + fun `inactive property returns all minus active`() = + runTest { + Given("entitlements with both active and inactive by product") { + val activeE = Entitlement("active", isActive = true) + val inactiveE = Entitlement("inactive_product", isActive = false) + val storage = + mockStorage( + storedProductEntitlements = + mapOf( + "prod1" to setOf(activeE), + "prod2" to setOf(inactiveE), + ), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(activeE))) + + When("accessing inactive property") { + val inactive = entitlements.inactive + + Then("it should contain the product entitlement not in active") { + assertTrue(inactive.any { it.id == "inactive_product" }) + } + And("it should not contain active entitlements") { + val activeIds = entitlements.active.map { it.id }.toSet() + assertTrue(inactive.none { it.id in activeIds }) + } + } + } + } + + @Test + fun `active property is empty when no sources have data`() = + runTest { + Given("a fresh Entitlements with no data") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + + Then("active should be empty") { + assertTrue(entitlements.active.isEmpty()) + } + } + } + + // ========================================== + // addEntitlementsByProductId + // ========================================== + + @Test + fun `addEntitlementsByProductId stores and makes entitlements available`() = + runTest { + Given("a fresh Entitlements instance") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val e1 = Entitlement("premium") + val e2 = Entitlement("basic") + val mapping = mapOf("prod_a" to setOf(e1), "prod_b" to setOf(e2)) + + When("adding entitlements by product ID") { + entitlements.addEntitlementsByProductId(mapping) + + Then("entitlementsByProductId should contain the mappings") { + assertEquals(setOf(e1), entitlements.entitlementsByProductId["prod_a"]) + assertEquals(setOf(e2), entitlements.entitlementsByProductId["prod_b"]) + } + And("all should include both entitlements") { + assertTrue(entitlements.all.contains(e1)) + assertTrue(entitlements.all.contains(e2)) + } + And("storage should be written to") { + verify { storage.write(StoredEntitlementsByProductId, any()) } + } + } + } + } + + @Test + fun `addEntitlementsByProductId overwrites existing product key`() = + runTest { + Given("existing entitlements for a product") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val oldE = Entitlement("old") + val newE = Entitlement("new") + + entitlements.addEntitlementsByProductId(mapOf("prod1" to setOf(oldE))) + + When("adding new entitlements for the same product") { + entitlements.addEntitlementsByProductId(mapOf("prod1" to setOf(newE))) + + Then("the new entitlement should replace the old one for that product") { + assertEquals(setOf(newE), entitlements.entitlementsByProductId["prod1"]) + } + And("all should reflect the update") { + assertTrue(entitlements.all.contains(newE)) + } + } + } + } + + @Test + fun `addEntitlementsByProductId with empty map does not crash`() = + runTest { + Given("a fresh Entitlements instance") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("adding an empty map") { + entitlements.addEntitlementsByProductId(emptyMap()) + + Then("entitlementsByProductId should remain empty") { + assertTrue(entitlements.entitlementsByProductId.isEmpty()) + } + } + } + } + + // ========================================== + // entitlementsByProductId Snapshot + // ========================================== + + @Test + fun `entitlementsByProductId returns a snapshot not a live reference`() = + runTest { + Given("entitlements with product mappings") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + entitlements.addEntitlementsByProductId(mapOf("prod1" to setOf(Entitlement("e1")))) + + When("taking a snapshot and then modifying the original") { + val snapshot = entitlements.entitlementsByProductId + entitlements.addEntitlementsByProductId(mapOf("prod2" to setOf(Entitlement("e2")))) + + Then("snapshot should not contain the new product") { + assertFalse(snapshot.containsKey("prod2")) + } + And("current entitlementsByProductId should contain both") { + assertTrue(entitlements.entitlementsByProductId.containsKey("prod1")) + assertTrue(entitlements.entitlementsByProductId.containsKey("prod2")) + } + } + } + } + + // ========================================== + // byProductId - Decomposed ID Matching + // ========================================== + + @Test + fun `byProductId exact match takes priority over partial`() = + runTest { + Given("entitlements mapped to both exact and partial matching product IDs") { + val exactE = Entitlement("exact_match") + val partialE = Entitlement("partial_match") + val storage = + mockStorage( + storedProductEntitlements = + mapOf( + "sub:plan:offer" to setOf(exactE), + "sub" to setOf(partialE), + ), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("querying with the exact full ID") { + val result = entitlements.byProductId("sub:plan:offer") + + Then("it should return the exact match entitlement") { + assertEquals(setOf(exactE), result) + } + } + } + } + + @Test + fun `byProductId falls back to subscriptionId contains match`() = + runTest { + Given("entitlements mapped only to a subscription ID") { + val e = Entitlement("sub_level") + val storage = + mockStorage( + storedProductEntitlements = mapOf("monthly_sub" to setOf(e)), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("querying with a full ID that contains the subscription ID") { + val result = entitlements.byProductId("monthly_sub:plan:offer") + + Then("it should fall back to contains match on subscriptionId") { + assertEquals(setOf(e), result) + } + } + } + } + + @Test + fun `byProductId returns empty for completely unknown product`() = + runTest { + Given("entitlements with some products") { + val storage = + mockStorage( + storedProductEntitlements = mapOf("known_product" to setOf(Entitlement("e"))), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("querying an unknown product") { + val result = entitlements.byProductId("completely_unknown") + + Then("result should be empty") { + assertTrue(result.isEmpty()) + } + } + } + } + + @Test + fun `byProductId skips products with empty entitlement sets`() = + runTest { + Given("a product mapped to an empty entitlement set") { + val fallbackE = Entitlement("fallback") + val storage = + mockStorage( + storedProductEntitlements = + mapOf( + "product_a" to emptySet(), + "product_a:plan" to setOf(fallbackE), + ), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("querying product_a") { + val result = entitlements.byProductId("product_a:plan") + + Then("it should skip the empty set and use the non-empty one") { + assertEquals(setOf(fallbackE), result) + } + } + } + } + + @Test + fun `byProductId simple product without colons`() = + runTest { + Given("a simple product ID with no base plan or offer") { + val e = Entitlement("simple") + val storage = + mockStorage( + storedProductEntitlements = mapOf("com.app.product" to setOf(e)), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("querying the simple product ID") { + val result = entitlements.byProductId("com.app.product") + + Then("it should find the exact match") { + assertEquals(setOf(e), result) + } + } + } + } + + // ========================================== + // byProductIds (batch) + // ========================================== + + @Test + fun `byProductIds returns union of entitlements from multiple products`() = + runTest { + Given("multiple products with different entitlements") { + val e1 = Entitlement("premium") + val e2 = Entitlement("addon") + val storage = + mockStorage( + storedProductEntitlements = + mapOf( + "prod1" to setOf(e1), + "prod2" to setOf(e2), + ), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("querying multiple product IDs") { + val result = entitlements.byProductIds(setOf("prod1", "prod2")) + + Then("result should contain entitlements from both products") { + assertTrue(result.contains(e1)) + assertTrue(result.contains(e2)) + assertEquals(2, result.size) + } + } + } + } + + @Test + fun `byProductIds with empty set returns empty`() = + runTest { + Given("entitlements with products") { + val storage = + mockStorage( + storedProductEntitlements = mapOf("prod1" to setOf(Entitlement("e"))), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("querying with empty set") { + val result = entitlements.byProductIds(emptySet()) + + Then("result should be empty") { + assertTrue(result.isEmpty()) + } + } + } + } + + @Test + fun `byProductIds deduplicates shared entitlements`() = + runTest { + Given("two products sharing the same entitlement") { + val shared = Entitlement("shared") + val storage = + mockStorage( + storedProductEntitlements = + mapOf( + "prod1" to setOf(shared), + "prod2" to setOf(shared), + ), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("querying both products") { + val result = entitlements.byProductIds(setOf("prod1", "prod2")) + + Then("result should contain the entitlement only once (set semantics)") { + assertEquals(1, result.size) + assertTrue(result.contains(shared)) + } + } + } + } + + @Test + fun `byProductIds with some unknown products returns only known`() = + runTest { + Given("one known and one unknown product") { + val e1 = Entitlement("known") + val storage = + mockStorage( + storedProductEntitlements = mapOf("known_prod" to setOf(e1)), + ) + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("querying both") { + val result = entitlements.byProductIds(setOf("known_prod", "unknown_prod")) + + Then("result should contain only the known entitlement") { + assertEquals(setOf(e1), result) + } + } + } + } + + // ========================================== + // Status Flow Persistence + // ========================================== + + @Test + fun `status changes are persisted to storage via flow collector`() = + runTest { + Given("Entitlements with backgroundScope for collector") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("setting Active status") { + val activeE = setOf(Entitlement("persisted")) + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(activeE)) + + Then("storage write should have been called with the new status") { + verify { + storage.write( + StoredSubscriptionStatus, + SubscriptionStatus.Active(activeE), + ) + } + } + } + } + } + + @Test + fun `Inactive status is persisted to storage`() = + runTest { + Given("Entitlements with backgroundScope") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("setting Inactive status") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) + + Then("Inactive should be persisted") { + verify { + storage.write(StoredSubscriptionStatus, SubscriptionStatus.Inactive) + } + } + } + } + } + + // ========================================== + // Web Entitlements Edge Cases + // ========================================== + + @Test + fun `web returns empty when redemption response has null customerInfo entitlements`() = + runTest { + Given("redemption response with no entitlements list") { + val redemption = + WebRedemptionResponse( + codes = emptyList(), + allCodes = emptyList(), + customerInfo = + CustomerInfo( + subscriptions = emptyList(), + nonSubscriptions = emptyList(), + userId = "user", + entitlements = emptyList(), + isPlaceholder = false, + ), + ) + val storage = mockStorage(redemptionResponse = redemption) + val entitlements = + makeEntitlements(storage, backgroundScope) + + Then("web should be empty") { + assertTrue(entitlements.web.isEmpty()) + } + } + } + + @Test + fun `web entitlements included in all property`() = + runTest { + Given("only web entitlements exist") { + val webE = Entitlement("web_only", isActive = true, store = Store.STRIPE) + val storage = mockStorage(redemptionResponse = webRedemption(webE)) + val entitlements = + makeEntitlements(storage, backgroundScope) + + Then("all should include web entitlements") { + assertTrue(entitlements.all.contains(webE)) + } + And("active should include web entitlements") { + assertTrue(entitlements.active.any { it.id == "web_only" }) + } + } + } + + @Test + fun `web entitlements in active even when status is Inactive`() = + runTest { + Given("Inactive status but web entitlements in storage") { + val webE = Entitlement("web_sub", isActive = true, store = Store.STRIPE) + val storage = mockStorage(redemptionResponse = webRedemption(webE)) + val entitlements = + makeEntitlements(storage, backgroundScope) + entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) + + Then("active should still contain web entitlements") { + assertTrue(entitlements.active.any { it.id == "web_sub" }) + } + } + } + + // ========================================== + // Deduplication / Merge Priority + // ========================================== + + @Test + fun `duplicate entitlement ID from status and device merges to single entry`() = + runTest { + Given("same entitlement ID from status and device sources") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val fromStatus = Entitlement("premium", isActive = true, store = Store.PLAY_STORE) + val fromDevice = Entitlement("premium", isActive = true, store = Store.PLAY_STORE) + + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(fromStatus))) + entitlements.activeDeviceEntitlements = setOf(fromDevice) + + When("accessing active") { + val active = entitlements.active + + Then("there should be only one premium entitlement after merge") { + assertEquals(1, active.count { it.id == "premium" }) + } + } + } + } + + @Test + fun `three sources with same ID deduplicate to one entry`() = + runTest { + Given("same entitlement ID from all three sources") { + val webE = Entitlement("premium", isActive = true, store = Store.STRIPE) + val storage = mockStorage(redemptionResponse = webRedemption(webE)) + val entitlements = + makeEntitlements(storage, backgroundScope) + + val statusE = Entitlement("premium", isActive = true, store = Store.PLAY_STORE) + val deviceE = Entitlement("premium", isActive = true) + + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusE))) + entitlements.activeDeviceEntitlements = setOf(deviceE) + + When("accessing active") { + val active = entitlements.active + + Then("only one premium entitlement should exist") { + assertEquals(1, active.count { it.id == "premium" }) + } + } + } + } + + // ========================================== + // Entitlement with Rich Properties + // ========================================== + + @Test + fun `entitlements with expiry dates and states are preserved through status transitions`() = + runTest { + Given("a richly-populated entitlement") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val now = Date() + val future = Date(now.time + 86400000) + val richE = + Entitlement( + id = "premium", + type = Entitlement.Type.SERVICE_LEVEL, + isActive = true, + productIds = setOf("prod_monthly", "prod_annual"), + latestProductId = "prod_annual", + startsAt = now, + renewedAt = now, + expiresAt = future, + isLifetime = false, + willRenew = true, + state = LatestSubscriptionState.SUBSCRIBED, + store = Store.PLAY_STORE, + ) + + When("setting Active with the rich entitlement") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(richE))) + + Then("active should contain the entitlement with all properties intact") { + val found = entitlements.active.first { it.id == "premium" } + assertEquals(setOf("prod_monthly", "prod_annual"), found.productIds) + assertEquals("prod_annual", found.latestProductId) + assertEquals(true, found.willRenew) + assertEquals(LatestSubscriptionState.SUBSCRIBED, found.state) + assertEquals(Store.PLAY_STORE, found.store) + assertEquals(future, found.expiresAt) + } + } + } + } + + // ========================================== + // Edge Cases + // ========================================== + + @Test + fun `setting Active with single entitlement works`() = + runTest { + Given("a single entitlement") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val e = Entitlement("solo") + + When("setting Active with single entitlement") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e))) + + Then("active should contain exactly one entitlement") { + assertEquals(1, entitlements.active.size) + assertTrue(entitlements.active.contains(e)) + } + } + } + } + + @Test + fun `setting Active with many entitlements works`() = + runTest { + Given("100 entitlements") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val many = (1..100).map { Entitlement("e_$it") }.toSet() + + When("setting Active with all of them") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(many)) + + Then("active should contain all 100") { + assertEquals(100, entitlements.active.size) + } + And("all should contain all 100") { + assertEquals(100, entitlements.all.size) + } + } + } + } + + @Test + fun `status flow value reflects latest status synchronously`() = + runTest { + Given("Entitlements instance") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + + When("setting status sequentially") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("a")))) + assertTrue(entitlements.status.value is SubscriptionStatus.Active) + + entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) + assertTrue(entitlements.status.value is SubscriptionStatus.Inactive) + + entitlements.setSubscriptionStatus(SubscriptionStatus.Unknown) + + Then("status value should match the latest set value") { + assertTrue(entitlements.status.value is SubscriptionStatus.Unknown) + } + } + } + } + + @Test + fun `addEntitlementsByProductId followed by byProductId returns correct result`() = + runTest { + Given("dynamically added product entitlements") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val e = Entitlement("dynamic") + + When("adding and then querying") { + entitlements.addEntitlementsByProductId(mapOf("dynamic_prod" to setOf(e))) + + Then("byProductId should find it") { + assertEquals(setOf(e), entitlements.byProductId("dynamic_prod")) + } + And("byProductIds should also find it") { + assertEquals(setOf(e), entitlements.byProductIds(setOf("dynamic_prod"))) + } + } + } + } + + @Test + fun `inactive returns empty when all are active`() = + runTest { + Given("only active entitlements") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val e = Entitlement("active", isActive = true) + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e))) + + Then("inactive should be empty") { + // inactive = _inactive + (all - active) + // all = {e}, active = {e}, so inactive additions = empty + assertTrue(entitlements.inactive.isEmpty()) + } + } + } + + @Test + fun `web property reflects setWebEntitlements updates`() = + runTest { + Given("entitlements with web entitlements set via actor") { + val webE1 = Entitlement("web_v1", isActive = true, store = Store.STRIPE) + val webE2 = Entitlement("web_v2", isActive = true, store = Store.STRIPE) + + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + + entitlements.setWebEntitlements(setOf(webE1)) + + When("first read returns v1") { + assertEquals(setOf(webE1), entitlements.web) + } + + entitlements.setWebEntitlements(setOf(webE2)) + + Then("second read should return v2") { + assertEquals(setOf(webE2), entitlements.web) + } + } + } + + @Test + fun `status collector can re-set the status with extra entitlements`() = + runTest { + Given("a status collector that adds a custom entitlement, like a delegate override") { + val storage = mockStorage() + val entitlements = makeEntitlements(storage, backgroundScope) + backgroundScope.launch(UnconfinedTestDispatcher(testScheduler)) { + entitlements.status.collect { status -> + if (status is SubscriptionStatus.Active && status.entitlements.none { it.id == "custom" }) { + entitlements.setSubscriptionStatus( + SubscriptionStatus.Active(status.entitlements + Entitlement("custom")), + ) + } + } + } + + When("the status becomes Active without the custom entitlement") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("pro")))) + advanceUntilIdle() + + Then("the status carries both entitlements") { + val status = entitlements.status.value as SubscriptionStatus.Active + assertEquals(setOf("pro", "custom"), status.entitlements.map { it.id }.toSet()) + } + And("both are active and the overridden status is persisted") { + assertEquals(setOf("pro", "custom"), entitlements.active.map { it.id }.toSet()) + verify { + storage.write( + StoredSubscriptionStatus, + match { + it is SubscriptionStatus.Active && + it.entitlements.map { e -> e.id }.toSet() == setOf("pro", "custom") + }, + ) + } + } + } + } + } + + @Test + fun `addEntitlementsByProductId drops entitlements of a remapped product from all`() = + runTest { + Given("a product mapped to an entitlement") { + val storage = mockStorage() + val entitlements = makeEntitlements(storage, backgroundScope) + entitlements.addEntitlementsByProductId(mapOf("p1" to setOf(Entitlement("pro")))) + + When("the same product is remapped to a different entitlement") { + entitlements.addEntitlementsByProductId(mapOf("p1" to setOf(Entitlement("premium")))) + + Then("all contains only the new entitlement") { + assertEquals(setOf("premium"), entitlements.all.map { it.id }.toSet()) + } + } + } + } + + @Test + fun `addEntitlementsByProductId accumulates entitlements across adds`() = + runTest { + Given("entitlements with existing product mappings") { + val storage = mockStorage() + val entitlements = + makeEntitlements(storage, backgroundScope) + val e1 = Entitlement("first") + val e2 = Entitlement("second") + + entitlements.addEntitlementsByProductId(mapOf("p1" to setOf(e1))) + + When("adding new product mappings (old key not overwritten)") { + entitlements.addEntitlementsByProductId(mapOf("p2" to setOf(e2))) + + Then("all should contain entitlements from both adds") { + assertTrue(entitlements.all.any { it.id == "first" }) + assertTrue(entitlements.all.any { it.id == "second" }) + } + } + } + } +} diff --git a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTest.kt b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTest.kt index 1a59f00e1..3e2fc75d9 100644 --- a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTest.kt @@ -18,6 +18,7 @@ import io.mockk.every import io.mockk.just import io.mockk.mockk import io.mockk.verify +import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.async import kotlinx.coroutines.delay @@ -49,10 +50,12 @@ class EntitlementsTest { Entitlement("test_entitlement"), ), ) - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("Entitlements is initialized") { - val entitlements = Entitlements(storage) + val entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) Then("it should load the stored status") { assertEquals(storedStatus, entitlements.status.value) @@ -81,7 +84,8 @@ class EntitlementsTest { } just Runs every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) When("setting active entitlement status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(activeEntitlements)) @@ -112,7 +116,8 @@ class EntitlementsTest { Given("an Entitlements instance") { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("setting active entitlement status with empty set") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(emptySet())) @@ -132,7 +137,8 @@ class EntitlementsTest { Given("an Entitlements instance") { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("setting NoActiveEntitlements status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) @@ -151,7 +157,8 @@ class EntitlementsTest { Given("an Entitlements instance") { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("setting Unknown status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Unknown) @@ -182,7 +189,8 @@ class EntitlementsTest { ), ) every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("creating a new Entitlements instance") { Then("it should return correct entitlements for each product") { @@ -216,7 +224,8 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("querying with subscription_monthly colon p1m colon freetrial") { val result = entitlements.byProductId("subscription_monthly:p1m:freetrial") @@ -243,7 +252,8 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("setting active device entitlements to only the active one") { entitlements.activeDeviceEntitlements = setOf(activeEntitlement) @@ -281,7 +291,8 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("no active device entitlements are set") { // activeDeviceEntitlements not set, should be empty @@ -313,7 +324,8 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) entitlements.activeDeviceEntitlements = setOf(activeEntitlement) When("subscription status is set to Inactive") { @@ -349,7 +361,8 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) When("setting both status and device entitlements") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusActiveEntitlement))) @@ -405,7 +418,8 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("accessing the web property") { - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) Then("it should return only active web entitlements") { assertEquals(setOf(webEntitlement1, webEntitlement2), entitlements.web) @@ -440,7 +454,8 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("accessing the web property") { - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) Then("it should return only active web entitlements") { assertEquals(setOf(activeWebEntitlement), entitlements.web) @@ -459,7 +474,8 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("accessing the web property") { - entitlements = Entitlements(storage) + entitlements = + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) Then("it should return empty set") { assertTrue(entitlements.web.isEmpty()) @@ -499,7 +515,8 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("setting subscription status (simulating external PC)") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusEntitlement))) Then("active should contain both status and web entitlements") { @@ -544,7 +561,8 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("external PC sets status with only its entitlements") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) // External PC sets status (like RC does) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(rcEntitlement))) @@ -587,7 +605,8 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) When("external PC reads web entitlements and merges them into status") { // This simulates what the updated RC controller does: @@ -644,7 +663,8 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(playEntitlement))) When("status is reset to Inactive (simulating sign out)") { @@ -690,7 +710,8 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) // Initial state entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("old_play")))) @@ -740,30 +761,17 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("userA_play")))) - When("user B identifies and storage is updated with user B's web entitlements") { + When("user B identifies and web entitlements are updated") { // Reset for user switch entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) - // Storage is updated with user B's web entitlements (simulating backend fetch) + // Web entitlements updated via actor (simulating WebPaywallRedeemer) val userBWebEntitlement = Entitlement("userB_web", isActive = true, store = Store.STRIPE) - val userBWebInfo = - CustomerInfo( - subscriptions = emptyList(), - nonSubscriptions = emptyList(), - userId = "userB", - entitlements = listOf(userBWebEntitlement), - isPlaceholder = false, - ) - val userBRedemption = - WebRedemptionResponse( - codes = emptyList(), - allCodes = emptyList(), - customerInfo = userBWebInfo, - ) - every { storage.read(LatestRedemptionResponse) } returns userBRedemption + entitlements.setWebEntitlements(setOf(userBWebEntitlement)) // User B's external PC sets status entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("userB_play")))) @@ -814,7 +822,8 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("all three sources have different entitlements") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusEntitlement))) entitlements.activeDeviceEntitlements = setOf(deviceEntitlement) @@ -857,7 +866,8 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("both sources have entitlement with same ID") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusPremium))) Then("active should deduplicate and contain only one premium entitlement") { @@ -894,7 +904,8 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("status is set to Unknown") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Unknown) Then("web property should still return web entitlements") { diff --git a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTestHelpers.kt b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTestHelpers.kt new file mode 100644 index 000000000..a9e7c8039 --- /dev/null +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTestHelpers.kt @@ -0,0 +1,16 @@ +package com.superwall.sdk.store + +import com.superwall.sdk.misc.primitives.StateActor +import com.superwall.sdk.storage.Storage +import kotlinx.coroutines.CoroutineScope + +/** Builds [Entitlements] the way [com.superwall.sdk.dependencies.DependencyContainer] does. */ +internal fun makeEntitlements( + storage: Storage, + scope: CoroutineScope, +): Entitlements = + Entitlements( + storage = storage, + actor = StateActor(createInitialEntitlementsState(storage), scope), + actorScope = scope, + ) diff --git a/superwall/src/test/java/com/superwall/sdk/store/testmode/TestModeTest.kt b/superwall/src/test/java/com/superwall/sdk/store/testmode/TestModeTest.kt index 2de682a97..202c27237 100644 --- a/superwall/src/test/java/com/superwall/sdk/store/testmode/TestModeTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/store/testmode/TestModeTest.kt @@ -941,7 +941,7 @@ class TestModeTest { val manager = makeManager() Then("initial state is Inactive") { - assertTrue(manager.state is TestModeState.Inactive) + assertTrue(manager.state.value is TestModeState.Inactive) assertFalse(manager.isTestMode) assertNull(manager.testModeReason) } @@ -957,7 +957,7 @@ class TestModeTest { } Then("state is Active with TestModeOption reason") { - val state = manager.state + val state = manager.state.value assertTrue(state is TestModeState.Active) assertEquals(TestModeReason.TestModeOption, (state as TestModeState.Active).reason) assertTrue(manager.isTestMode) @@ -969,7 +969,7 @@ class TestModeTest { } Then("state is back to Inactive") { - assertTrue(manager.state is TestModeState.Inactive) + assertTrue(manager.state.value is TestModeState.Inactive) assertFalse(manager.isTestMode) assertNull(manager.testModeReason) } @@ -1105,4 +1105,136 @@ class TestModeTest { } // endregion + + // region presentModal — UI flow + + @Test + fun `activate with justActivated=true and no activity falls back to default subscription status`() = + kotlinx.coroutines.test.runTest { + val storage = makeStorage() + val entitlements = mockk(relaxed = true) + val tracked = mutableListOf() + val manager = + TestMode( + storage = storage, + isTestEnvironment = false, + getSuperwallProducts = { + com.superwall.sdk.misc.Either.Success( + com.superwall.sdk.store.testmode.models.SuperwallProductsResponse(data = emptyList()), + ) + }, + entitlements = entitlements, + activityProvider = { null }, + activityTracker = { null }, + tracker = { tracked.add(it) }, + ) + activateTestMode(manager) + + manager.activate(makeConfig(), justActivated = true) + + // No Open/Close tracking when modal could not be presented. + assertTrue( + "TestModeModal Open/Close must not be tracked when no activity is available", + tracked.none { it is com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent.TestModeModal }, + ) + // Fallback path sets the default (empty entitlements → Inactive) status + // both on TestMode itself and on the entitlements collaborator. + assertEquals(SubscriptionStatus.Inactive, manager.overriddenSubscriptionStatus) + verify(exactly = 1) { entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) } + } + + @Test + fun `activate with justActivated=true wires modal result into settings, status, entitlements, and tracking`() = + kotlinx.coroutines.test.runTest { + val storage = makeStorage() + // Relaxed mockk's generic `read` defaults to `Any`, which can't cast to TestModeSettings. + every { storage.read(com.superwall.sdk.storage.StoredTestModeSettings) } returns null + val entitlements = mockk(relaxed = true) + val tracked = mutableListOf() + val activity = mockk(relaxed = true) + val activityProvider = mockk(relaxed = true).also { + every { it.getCurrentActivity() } returns activity + } + val modalResult = + com.superwall.sdk.store.testmode.ui.TestModeModalResult( + entitlements = + listOf( + com.superwall.sdk.store.testmode.ui.EntitlementSelection( + identifier = "pro", + state = com.superwall.sdk.store.testmode.ui.EntitlementStateOption.Subscribed, + ), + ), + freeTrialOverride = FreeTrialOverride.ForceAvailable, + ) + val capturedSavedSettings = mutableListOf() + val manager = + TestMode( + storage = storage, + isTestEnvironment = false, + getSuperwallProducts = { + com.superwall.sdk.misc.Either.Success( + com.superwall.sdk.store.testmode.models.SuperwallProductsResponse(data = emptyList()), + ) + }, + entitlements = entitlements, + activityProvider = { activityProvider }, + activityTracker = { null }, + apiKey = { "test-api-key" }, + dashboardBaseUrl = { "https://dash" }, + tracker = { tracked.add(it) }, + showModal = { _, _, _, _, _, _, savedSettings -> + capturedSavedSettings.add(savedSettings) + modalResult + }, + ) + activateTestMode(manager) + + manager.activate(makeConfig(), justActivated = true) + + // Free-trial override + entitlement selections from the modal are applied. + assertEquals(FreeTrialOverride.ForceAvailable, manager.freeTrialOverride) + assertEquals( + listOf("pro"), + manager.testEntitlementSelections.map { it.identifier }, + ) + assertEquals(setOf("pro"), manager.testEntitlementIds) + + // Settings are persisted with the same selections + override. + verify { + storage.write( + com.superwall.sdk.storage.StoredTestModeSettings, + match { + it.freeTrialOverride == FreeTrialOverride.ForceAvailable && + it.entitlementSelections.map { sel -> sel.identifier } == listOf("pro") + }, + ) + } + + // Subscription status reflects the active selection on both + // TestMode itself and the entitlements collaborator. + val status = manager.overriddenSubscriptionStatus + assertTrue("Expected SubscriptionStatus.Active", status is SubscriptionStatus.Active) + val granted = (status as SubscriptionStatus.Active).entitlements + assertEquals(setOf("pro"), granted.map { it.id }.toSet()) + assertTrue("Subscribed selection must grant an active entitlement", granted.all { it.isActive }) + verify(exactly = 1) { entitlements.setSubscriptionStatus(status) } + + // Open is tracked before showModal, Close after — verify both the + // emission and ordering. + val modalEvents = + tracked.filterIsInstance() + assertEquals( + listOf( + com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent.TestModeModal.State.Open, + com.superwall.sdk.analytics.internal.trackable.InternalSuperwallEvent.TestModeModal.State.Close, + ), + modalEvents.map { it.state }, + ) + + // savedSettings forwarded to the modal launcher (null on first run — + // storage mock returns null for read). + assertEquals(1, capturedSavedSettings.size) + } + + // endregion } diff --git a/superwall/src/test/java/com/superwall/sdk/web/WebPaywallRedeemerTest.kt b/superwall/src/test/java/com/superwall/sdk/web/WebPaywallRedeemerTest.kt index 142432378..41460a265 100644 --- a/superwall/src/test/java/com/superwall/sdk/web/WebPaywallRedeemerTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/web/WebPaywallRedeemerTest.kt @@ -1,6 +1,7 @@ package com.superwall.sdk.web import android.content.Context +import com.superwall.sdk.And import com.superwall.sdk.Given import com.superwall.sdk.Then import com.superwall.sdk.When @@ -42,6 +43,8 @@ import kotlinx.serialization.json.JsonArray import kotlinx.serialization.json.JsonElement import kotlinx.serialization.json.JsonPrimitive import kotlinx.serialization.json.buildJsonObject +import org.junit.Assert.assertEquals +import org.junit.Assert.assertTrue import org.junit.Before import org.junit.Test @@ -131,6 +134,7 @@ class WebPaywallRedeemerTest { ) }, var getIntegrationPropsFn: () -> Map = { emptyMap() }, + var setWebEntitlementsFn: (Set) -> Unit = {}, ) : WebPaywallRedeemer.Factory { override fun willRedeemLink() = willRedeemLinkFn() @@ -150,6 +154,8 @@ class WebPaywallRedeemerTest { override fun internallySetSubscriptionStatus(status: SubscriptionStatus) = this@WebPaywallRedeemerTest.setSubscriptionStatus(status) + override fun setWebEntitlements(entitlements: Set) = setWebEntitlementsFn(entitlements) + override suspend fun isPaywallVisible(): Boolean = this@WebPaywallRedeemerTest.isPaywallVisible() override suspend fun triggerRestoreInPaywall() = this@WebPaywallRedeemerTest.showRestoreDialogAndDismiss() @@ -239,6 +245,8 @@ class WebPaywallRedeemerTest { ) } returns Either.Success(response) + val published = java.util.concurrent.CopyOnWriteArrayList>() + When("creating redeemer and advancing scheduler") { redeemer = WebPaywallRedeemer( @@ -248,7 +256,7 @@ class WebPaywallRedeemerTest { network, storage, customerInfoManager = mockk(relaxed = true), - factory = TestFactory(), + factory = TestFactory(setWebEntitlementsFn = { published.add(it) }), ) testScheduler.advanceUntilIdle() @@ -256,9 +264,14 @@ class WebPaywallRedeemerTest { verify(exactly = 1) { storage.write(LatestRedemptionResponse, response) } - println(mutableEntitlements) assert(mutableEntitlements == setOf(webEntitlement, normalEntitlement)) } + + And("it publishes exactly the web entitlements it persisted") { + // Polling also runs, but storage holds no redemption response in this + // mock, so it must not publish entitlements it cannot persist. + assertEquals(listOf(setOf(webEntitlement)), published.toList()) + } } } } @@ -926,4 +939,85 @@ class WebPaywallRedeemerTest { } } } + + @Test + fun `user switch clears web entitlements in Entitlements through the factory`() { + Given("user A's web redemption is stored and restored into Entitlements on start") { + val userAWeb = Entitlement("userA_web", isActive = true) + val userAResponse = + WebRedemptionResponse( + codes = + listOf( + RedemptionResult.Success( + code = "userA_code", + redemptionInfo = + RedemptionInfo( + ownership = RedemptionOwnership.AppUser(appUserId = "userA"), + purchaserInfo = + PurchaserInfo( + "userA", + email = null, + storeIdentifiers = StoreIdentifiers.Stripe("123", emptyList()), + ), + entitlements = listOf(userAWeb), + ), + ), + ), + customerInfo = + CustomerInfo( + subscriptions = emptyList(), + nonSubscriptions = emptyList(), + userId = "userA", + entitlements = listOf(userAWeb), + isPlaceholder = false, + ), + ) + val storage = + object : Storage { + val values = mutableMapOf() + + @Suppress("UNCHECKED_CAST") + override fun read(storable: Storable): T? = values[storable.key] as T? + + override fun write( + storable: Storable, + data: T, + ) { + values[storable.key] = data + } + + override fun delete(storable: Storable) { + values.remove(storable.key) + } + + override fun clean() = values.clear() + } + storage.write(LatestRedemptionResponse, userAResponse) + + val entitlementsScope = kotlinx.coroutines.CoroutineScope(kotlinx.coroutines.Dispatchers.Unconfined) + val entitlements = com.superwall.sdk.store.makeEntitlements(storage, entitlementsScope) + // Wire the redeemer to Entitlements the way DependencyContainer.setWebEntitlements does. + redeemer = + WebPaywallRedeemer( + context, + IOScope(testDispatcher), + deepLinkReferrer, + network, + storage, + customerInfoManager = mockk(relaxed = true), + factory = TestFactory(setWebEntitlementsFn = { entitlements.setWebEntitlements(it) }), + ) + assertEquals(setOf(userAWeb), entitlements.web) + + When("Superwall.reset wipes storage and then clears the user's redemptions") { + storage.clean() + redeemer.clear(RedemptionOwnershipType.AppUser) + + Then("user A's web entitlements are gone from Entitlements") { + assertEquals(emptySet(), entitlements.web) + assertTrue(entitlements.active.none { it.id == "userA_web" }) + } + } + } + } }