From a61b38610fbb7076154ecd7ffa292f376ac975a4 Mon Sep 17 00:00:00 2001 From: Ian Rumac Date: Sat, 2 May 2026 13:15:35 +0200 Subject: [PATCH 01/10] Add actors for PV cache and TestMode Migrates PaywallViewCache and TestMode onto the StateActor primitives: PaywallCacheState/PaywallCacheContext own the view map and active key, TestModeState gains pure reducers and async actions with TestModeContext and TestModeLogic, and ConfigState moves up to the config package. Rebased onto develop (Sep 2026) with these adaptations: - Keep develop's loadingColor on PaywallViewCache/LoadingView. - Port develop's productsLoaded/awaitTestProducts so StoreManager can wait for the test product catalog; the deferred is carried through session copies. - SetActive preserves an existing session and only clears entitlement selections when the activation reason changes, matching develop. - TestMode.activate is suspend and awaited via `immediate`; ConfigState launches it in its scope so the modal never blocks config. - Cache view factories no longer hop to Dispatchers.Main: acquire* block the caller (usually main) with runBlocking, so that hop deadlocked. - Keep develop's ensureActive guard in PaywallMessageHandler. Co-Authored-By: Claude Fable 5.1 --- .../config/ConfigManagerInstrumentedTest.kt | 2 - .../main/java/com/superwall/sdk/Superwall.kt | 2 +- .../sdk/analytics/internal/Tracking.kt | 2 +- .../analytics/session/AppSessionManager.kt | 2 +- .../com/superwall/sdk/config/ConfigContext.kt | 4 - .../com/superwall/sdk/config/ConfigManager.kt | 5 - .../sdk/config/{models => }/ConfigState.kt | 10 +- .../sdk/dependencies/DependencyContainer.kt | 12 +- .../sdk/misc/Config+AwaitFirstValidConfig.kt | 2 +- .../sdk/paywall/manager/PaywallManager.kt | 2 +- .../sdk/paywall/manager/PaywallViewCache.kt | 253 +++++++--- .../operators/WaitForSubsStatusAndConfig.kt | 3 +- .../view/webview/templating/TemplateLogic.kt | 1 + .../superwall/sdk/store/testmode/TestMode.kt | 385 +++++---------- .../sdk/store/testmode/TestModeContext.kt | 39 ++ .../sdk/store/testmode/TestModeLogic.kt | 64 +++ .../sdk/store/testmode/TestModeState.kt | 214 ++++++++- .../sdk/store/testmode/ui/TestModeModal.kt | 2 +- .../com/superwall/sdk/SdkContextImplTest.kt | 2 +- .../superwall/sdk/config/ConfigManagerTest.kt | 60 ++- .../sdk/config/ConfigStateReducerTest.kt | 1 - .../sdk/config/PaywallPreloadTest.kt | 8 +- .../sdk/misc/AwaitFirstValidConfigTest.kt | 2 +- .../sdk/paywall/manager/PaywallManagerTest.kt | 28 +- .../paywall/manager/PaywallViewCacheTest.kt | 441 ++++++++++++++++++ .../WaitForSubsStatusAndConfigTest.kt | 2 +- .../sdk/store/testmode/TestModeTest.kt | 139 +++++- 27 files changed, 1281 insertions(+), 406 deletions(-) rename superwall/src/main/java/com/superwall/sdk/config/{models => }/ConfigState.kt (98%) create mode 100644 superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeContext.kt create mode 100644 superwall/src/main/java/com/superwall/sdk/store/testmode/TestModeLogic.kt create mode 100644 superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt 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/main/java/com/superwall/sdk/Superwall.kt b/superwall/src/main/java/com/superwall/sdk/Superwall.kt index cdcc333c4..b8b76fd7a 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 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/dependencies/DependencyContainer.kt b/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt index da7d94c79..245169c61 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,7 +53,6 @@ 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 @@ -295,7 +295,8 @@ class DependencyContainer( else -> "https://superwall.com" } }, - track = { Superwall.instance.track(it) }, + tracker = { Superwall.instance.track(it) }, + ioScope = ioScope, ) testModeTransactionHandler = TestModeTransactionHandler( @@ -463,8 +464,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 +493,6 @@ class DependencyContainer( setSubscriptionStatus = { status -> entitlements.setSubscriptionStatus(status) }, - activateTestMode = { config, justActivated -> - testMode.activate(config, justActivated) - }, actor = configActor, ) 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..4ed66f7d6 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 @@ -44,7 +44,7 @@ class PaywallManager( 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..c0d33e646 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,222 @@ 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.Dispatchers +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) }) + + object RemoveAllExceptActive : Updates({ state -> + val active = state.activePaywallVcKey + val kept = if (active != null) state.views.filterKeys { it == active } else emptyMap() + state.copy(views = kept) + }) + + 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)) + }) + + object RemoveAllExceptActive : Actions({ + val active = state.value.activePaywallVcKey + state.value.views.keys + .filter { it != active } + .forEach { viewStorage.removeView(it) } + update(Updates.RemoveAllExceptActive) + }) + + /** + * 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) access it directly. + */ +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 = + SequentialActor(PaywallCacheState(), CoroutineScope(Dispatchers.IO)), +) : 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 getAllPaywallViews(): List = state.value.paywallViews - fun save( + 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 - } + suspend fun removeAll() { + immediate(PaywallCacheState.Actions.RemoveAllExceptActive) } - fun getPaywallView(key: String): PaywallView? = - try { - store.retrieveView(key) as PaywallView? - } catch (e: Throwable) { - null + /** + * 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 removePaywallView(identifier: PaywallIdentifier) { - store.removeView( - PaywallCacheLogic.key( - identifier, - locale = deviceHelper.locale, - ), - ) } - fun removeAll() { - store.views.keys.forEach { key -> - if (key != _activePaywallVcKey) { - store.removeView(key) - } + 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/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/webview/templating/TemplateLogic.kt b/superwall/src/main/java/com/superwall/sdk/paywall/view/webview/templating/TemplateLogic.kt index 9674b44fb..1ec652ebd 100644 --- a/superwall/src/main/java/com/superwall/sdk/paywall/view/webview/templating/TemplateLogic.kt +++ b/superwall/src/main/java/com/superwall/sdk/paywall/view/webview/templating/TemplateLogic.kt @@ -14,6 +14,7 @@ import com.superwall.sdk.paywall.view.webview.templating.models.ProductTemplate import kotlinx.serialization.json.Json object TemplateLogic { + suspend fun getBase64EncodedTemplates( json: Json, paywall: Paywall, 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/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/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/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..aabd96476 --- /dev/null +++ b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt @@ -0,0 +1,441 @@ +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.network.device.DeviceHelper +import com.superwall.sdk.paywall.view.LoadingView +import com.superwall.sdk.paywall.view.PaywallView +import com.superwall.sdk.paywall.view.ViewStorage +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.async +import kotlinx.coroutines.awaitAll +import kotlinx.coroutines.launch +import kotlinx.coroutines.test.runTest +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 fun newCache(): PaywallViewCache = + PaywallViewCache(appCtx, storage, activityProvider, deviceHelper) + + @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() + } + } + + // ------------------------------------------------------------------- + // 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 + assertNotNull(final) + assertTrue(final!!.startsWith("k_")) + } + } + } + + @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("the cache does not crash and getAllPaywallViews is consistent") { + val views = cache.getAllPaywallViews() + // Final state depends on interleaving but must not throw + assertTrue(views.size in 0..ids.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"))) + } + } + } + } +} 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/testmode/TestModeTest.kt b/superwall/src/test/java/com/superwall/sdk/store/testmode/TestModeTest.kt index 2de682a97..632efb438 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,137 @@ 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 expectedStatus = manager.buildSubscriptionStatus() + assertTrue( + "Expected SubscriptionStatus.Active", + expectedStatus is SubscriptionStatus.Active, + ) + assertEquals(expectedStatus, manager.overriddenSubscriptionStatus) + verify(exactly = 1) { entitlements.setSubscriptionStatus(expectedStatus) } + + // 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 } From ff0339d547caf961e62cd57cf39c7a7bd1c104e3 Mon Sep 17 00:00:00 2001 From: Ian Rumac Date: Tue, 15 Sep 2026 14:52:02 +0200 Subject: [PATCH 02/10] Move Entitlements onto the StateActor primitives Ports the entitlements slice of the March actor draft (ir/refactor/actors) onto current develop: - EntitlementsState holds status, product entitlements, device, backing and web entitlements as an immutable snapshot with pure reducers; `all`, `active` and `inactive` are derived from it. createInitialEntitlementsState rebuilds the snapshot from storage before the actor starts, including web entitlements from the latest redemption response. - Entitlements becomes a facade over a StateActor implementing EntitlementsContext. Status changes and product entitlement updates are persisted immediately instead of through a collected flow. - Web entitlements are cached in state rather than re-read from storage on every access, so WebPaywallRedeemer now publishes them through a new Factory.setWebEntitlements hook at each point it writes the redemption response. - Adds EntitlementsRefactorSafetyTest (46 cases) and reworks EntitlementsTest for the actor construction. Differences from the draft: the constructor keeps `Entitlements(storage)` working via defaults, and EntitlementsContext no longer carries HasExternalPurchaseControllerFactory since no action used it. Co-Authored-By: Claude Fable 5.1 --- .../sdk/dependencies/DependencyContainer.kt | 7 +- .../com/superwall/sdk/store/Entitlements.kt | 207 +-- .../sdk/store/EntitlementsContext.kt | 11 + .../superwall/sdk/store/EntitlementsState.kt | 199 ++ .../superwall/sdk/web/WebPaywallRedeemer.kt | 19 + .../store/EntitlementsRefactorSafetyTest.kt | 1640 +++++++++++++++++ .../superwall/sdk/store/EntitlementsTest.kt | 274 ++- .../sdk/web/WebPaywallRedeemerTest.kt | 2 + 8 files changed, 2174 insertions(+), 185 deletions(-) create mode 100644 superwall/src/main/java/com/superwall/sdk/store/EntitlementsContext.kt create mode 100644 superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt create mode 100644 superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt 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 245169c61..b9f9074c5 100644 --- a/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt +++ b/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt @@ -57,6 +57,7 @@ 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 @@ -278,7 +279,7 @@ class DependencyContainer( json = json(), _apiKey = apiKey ) - entitlements = Entitlements(storage) + entitlements = Entitlements(storage, actorScope = ioScope) val options = options ?: SuperwallOptions() testMode = TestMode( @@ -1260,6 +1261,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/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..ed1f94450 --- /dev/null +++ b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt @@ -0,0 +1,199 @@ +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 + product 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, + allTracked = newProducts.values.flatten().toSet(), + ) + }) + + 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() + + // 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 + } + } + + if (productEntitlements != null) { + state = EntitlementsState.Updates.AddProductEntitlements(productEntitlements).reduce(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/web/WebPaywallRedeemer.kt b/superwall/src/main/java/com/superwall/sdk/web/WebPaywallRedeemer.kt index d9b5d6e1b..24332174d 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( @@ -523,6 +539,9 @@ class WebPaywallRedeemer( updatedResponse, ) } + factory.setWebEntitlements( + newEntitlements.filter { it.isActive }.toSet(), + ) // Trigger CustomerInfo merge customerInfoManager.updateMergedCustomerInfo() 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..91c5a234c --- /dev/null +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt @@ -0,0 +1,1640 @@ +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.misc.primitives.StateActor +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.Dispatchers +import kotlinx.coroutines.async +import kotlinx.coroutines.delay +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 +import kotlin.time.Duration.Companion.seconds + +/** + * 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 all inactive entitlements becomes Inactive`() = + runTest { + Given("entitlements that are all inactive") { + val storage = mockStorage() + val entitlements = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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") { + // active entitlement may appear in inactive if the exact object differs + // but we check the concept + val activeIds = entitlements.active.map { it.id }.toSet() + val purelyInactive = inactive.filter { it.id !in activeIds } + assertTrue(purelyInactive.any { it.id == "inactive_product" }) + } + } + } + } + + @Test + fun `active property is empty when no sources have data`() = + runTest { + Given("a fresh Entitlements with no data") { + val storage = mockStorage() + val entitlements = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } + + When("setting Active status") { + val activeE = setOf(Entitlement("persisted")) + entitlements.setSubscriptionStatus(SubscriptionStatus.Active(activeE)) + + // Give collector time to process + async(Dispatchers.Default) { delay(1.seconds) }.await() + + 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } + + When("setting Inactive status") { + entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) + async(Dispatchers.Default) { delay(1.seconds) }.await() + + 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 `addEntitlementsByProductId clears and rebuilds _all`() = + runTest { + Given("entitlements with existing product mappings") { + val storage = mockStorage() + val entitlements = + StateActor( + createInitialEntitlementsState(storage), + backgroundScope, + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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..69dabd574 100644 --- a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTest.kt @@ -4,6 +4,7 @@ import com.superwall.sdk.And import com.superwall.sdk.Given import com.superwall.sdk.Then import com.superwall.sdk.When +import com.superwall.sdk.misc.primitives.StateActor import com.superwall.sdk.models.customer.CustomerInfo import com.superwall.sdk.models.entitlements.Entitlement import com.superwall.sdk.models.entitlements.SubscriptionStatus @@ -18,6 +19,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 +51,30 @@ class EntitlementsTest { Entitlement("test_entitlement"), ), ) - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } When("Entitlements is initialized") { - val entitlements = Entitlements(storage) + val entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } Then("it should load the stored status") { assertEquals(storedStatus, entitlements.status.value) @@ -81,7 +103,14 @@ class EntitlementsTest { } just Runs every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } When("setting active entitlement status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(activeEntitlements)) @@ -112,7 +141,17 @@ class EntitlementsTest { Given("an Entitlements instance") { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } When("setting active entitlement status with empty set") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(emptySet())) @@ -132,7 +171,17 @@ class EntitlementsTest { Given("an Entitlements instance") { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } When("setting NoActiveEntitlements status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) @@ -151,7 +200,17 @@ class EntitlementsTest { Given("an Entitlements instance") { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } When("setting Unknown status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Unknown) @@ -182,7 +241,17 @@ class EntitlementsTest { ), ) every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } When("creating a new Entitlements instance") { Then("it should return correct entitlements for each product") { @@ -216,7 +285,17 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } When("querying with subscription_monthly colon p1m colon freetrial") { val result = entitlements.byProductId("subscription_monthly:p1m:freetrial") @@ -243,7 +322,17 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } When("setting active device entitlements to only the active one") { entitlements.activeDeviceEntitlements = setOf(activeEntitlement) @@ -281,7 +370,17 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } When("no active device entitlements are set") { // activeDeviceEntitlements not set, should be empty @@ -313,7 +412,14 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } entitlements.activeDeviceEntitlements = setOf(activeEntitlement) When("subscription status is set to Inactive") { @@ -349,7 +455,14 @@ class EntitlementsTest { ) every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } When("setting both status and device entitlements") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusActiveEntitlement))) @@ -405,7 +518,17 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("accessing the web property") { - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } Then("it should return only active web entitlements") { assertEquals(setOf(webEntitlement1, webEntitlement2), entitlements.web) @@ -440,7 +563,17 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("accessing the web property") { - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } Then("it should return only active web entitlements") { assertEquals(setOf(activeWebEntitlement), entitlements.web) @@ -459,7 +592,17 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("accessing the web property") { - entitlements = Entitlements(storage) + entitlements = + StateActor( + createInitialEntitlementsState(storage), + CoroutineScope(Dispatchers.Default), + ).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = actor.scope, + ) + } Then("it should return empty set") { assertTrue(entitlements.web.isEmpty()) @@ -499,7 +642,14 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("setting subscription status (simulating external PC)") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusEntitlement))) Then("active should contain both status and web entitlements") { @@ -544,7 +694,14 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("external PC sets status with only its entitlements") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } // External PC sets status (like RC does) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(rcEntitlement))) @@ -587,7 +744,14 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } When("external PC reads web entitlements and merges them into status") { // This simulates what the updated RC controller does: @@ -644,7 +808,14 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(playEntitlement))) When("status is reset to Inactive (simulating sign out)") { @@ -690,7 +861,14 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } // Initial state entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("old_play")))) @@ -740,30 +918,23 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = 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 +985,14 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("all three sources have different entitlements") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusEntitlement))) entitlements.activeDeviceEntitlements = setOf(deviceEntitlement) @@ -857,7 +1035,14 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("both sources have entitlement with same ID") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusPremium))) Then("active should deduplicate and contain only one premium entitlement") { @@ -894,7 +1079,14 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null When("status is set to Unknown") { - entitlements = Entitlements(storage, scope = backgroundScope) + entitlements = + StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> + Entitlements( + storage = storage, + actor = actor, + actorScope = backgroundScope, + ) + } entitlements.setSubscriptionStatus(SubscriptionStatus.Unknown) Then("web property should still return web entitlements") { 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..36e93d35d 100644 --- a/superwall/src/test/java/com/superwall/sdk/web/WebPaywallRedeemerTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/web/WebPaywallRedeemerTest.kt @@ -150,6 +150,8 @@ class WebPaywallRedeemerTest { override fun internallySetSubscriptionStatus(status: SubscriptionStatus) = this@WebPaywallRedeemerTest.setSubscriptionStatus(status) + override fun setWebEntitlements(entitlements: Set) {} + override suspend fun isPaywallVisible(): Boolean = this@WebPaywallRedeemerTest.isPaywallVisible() override suspend fun triggerRestoreInPaywall() = this@WebPaywallRedeemerTest.showRestoreDialogAndDismiss() From fa17c96dbc4d72cc53077826687d597c51680c78 Mon Sep 17 00:00:00 2001 From: Ian Rumac Date: Fri, 25 Sep 2026 13:22:16 +0200 Subject: [PATCH 03/10] Address review on actor refactor - Acquire loading/shimmer views through the cache in SuperwallPaywallActivity so getPaywall() + startWithView() no longer crashes when they have not been created yet. - Route external ViewStorage writes (activity launch/restore, DebugView) through new synchronous PaywallViewCache.storeView/removeView so cache state and ViewStorage cannot diverge; RemoveAllExceptActive now removes exactly the keys it evicted. - Run the cache actor on the container ioScope instead of a leaked scope. - Only publish polled web entitlements when they are persisted. - Tests: cover external writers and lazy loading/shimmer, assert the redeemer publishes what it persists, tighten weak assertions, drop obsolete delays, and extract a makeEntitlements helper. Co-Authored-By: Claude Opus 5.5 --- .../java/com/superwall/sdk/debug/DebugView.kt | 3 +- .../sdk/dependencies/DependencyContainer.kt | 2 + .../sdk/paywall/manager/PaywallManager.kt | 6 +- .../sdk/paywall/manager/PaywallViewCache.kt | 50 +- .../paywall/view/SuperwallPaywallActivity.kt | 19 +- .../view/webview/templating/TemplateLogic.kt | 1 - .../superwall/sdk/web/WebPaywallRedeemer.kt | 8 +- .../paywall/manager/PaywallViewCacheTest.kt | 137 ++++- .../store/EntitlementsRefactorSafetyTest.kt | 522 ++---------------- .../superwall/sdk/store/EntitlementsTest.kt | 229 +------- .../sdk/store/EntitlementsTestHelpers.kt | 16 + .../sdk/store/testmode/TestModeTest.kt | 13 +- .../sdk/web/WebPaywallRedeemerTest.kt | 16 +- 13 files changed, 296 insertions(+), 726 deletions(-) create mode 100644 superwall/src/test/java/com/superwall/sdk/store/EntitlementsTestHelpers.kt 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..9ca50698a 100644 --- a/superwall/src/main/java/com/superwall/sdk/debug/DebugView.kt +++ b/superwall/src/main/java/com/superwall/sdk/debug/DebugView.kt @@ -923,8 +923,7 @@ internal class DebugViewActivity : AppCompatActivity() { view: View, ) { val key = UUID.randomUUID().toString() - Superwall.instance.dependencyContainer - .makeViewStore() + Superwall.instance.dependencyContainer.paywallManager.cache .storeView(key, view) val intent = 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 b9f9074c5..3c2654faf 100644 --- a/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt +++ b/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt @@ -76,6 +76,7 @@ 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.presentation.CustomCallbackRegistry @@ -908,6 +909,7 @@ class DependencyContainer( activityProvider!!, deviceHelper, configManager.options.paywalls.loadingColor, + actor = SequentialActor(PaywallCacheState(), ioScope), ) override fun activePaywallId(): String? = 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 4ed66f7d6..d75e20fb7 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 @@ -30,7 +30,11 @@ class PaywallManager( private var _cache: PaywallViewCache? = null - private val cache: PaywallViewCache + /** + * The single cache instance. Exposed so Activities that write to + * [com.superwall.sdk.paywall.view.ViewStorage] can go through it instead. + */ + internal val cache: PaywallViewCache get() { if (_cache == null) { _cache = createCache() 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 c0d33e646..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 @@ -17,7 +17,6 @@ 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.Dispatchers import kotlinx.coroutines.runBlocking /** @@ -55,11 +54,9 @@ data class PaywallCacheState( val key: String?, ) : Updates({ it.copy(activePaywallVcKey = key) }) - object RemoveAllExceptActive : Updates({ state -> - val active = state.activePaywallVcKey - val kept = if (active != null) state.views.filterKeys { it == active } else emptyMap() - state.copy(views = kept) - }) + data class RemoveViews( + val keys: Set, + ) : Updates({ it.copy(views = it.views - keys) }) data class Hydrate( val views: Map, @@ -87,12 +84,16 @@ data class PaywallCacheState( 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 - state.value.views.keys - .filter { it != active } - .forEach { viewStorage.removeView(it) } - update(Updates.RemoveAllExceptActive) + val evicted = state.value.views.keys.filterTo(mutableSetOf()) { it != active } + evicted.forEach { viewStorage.removeView(it) } + update(Updates.RemoveViews(evicted)) }) /** @@ -100,7 +101,7 @@ data class PaywallCacheState( * 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`, + * 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( @@ -148,7 +149,9 @@ interface PaywallCacheContext : StoreContext = - SequentialActor(PaywallCacheState(), CoroutineScope(Dispatchers.IO)), + override val actor: SequentialActor, ) : PaywallCacheContext { override val scope: CoroutineScope get() = actor.scope @@ -197,6 +199,26 @@ class PaywallViewCache( immediate(PaywallCacheState.Actions.Remove(identifier)) } + /** + * 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)) + } + + /** Synchronous counterpart of [storeView]. */ + fun removeView(key: String) { + viewStorage.removeView(key) + actor.update(PaywallCacheState.Updates.RemoveView(key)) + } + suspend fun removeAll() { immediate(PaywallCacheState.Actions.RemoveAllExceptActive) } 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..f7601db0b 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 @@ -157,8 +157,7 @@ class SuperwallPaywallActivity : AppCompatActivity() { } return launchPaywallActivity(context, intent).onFailure { - Superwall.instance.dependencyContainer - .makeViewStore() + Superwall.instance.dependencyContainer.paywallManager.cache .removeView(key) view.clearActivityLaunchState() } @@ -167,12 +166,13 @@ class SuperwallPaywallActivity : AppCompatActivity() { private fun PaywallView.prepareViewForDisplay(key: String) { webView.enableBackgroundRendering() webView.attach(this) - val viewStorageViewModel = Superwall.instance.dependencyContainer.makeViewStore() + val cache = Superwall.instance.dependencyContainer.paywallManager.cache // 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 = cache.acquireLoadingView() val style = state.paywall.presentation.style val shimmer = if (style is PaywallPresentationStyle.Popup) { @@ -184,12 +184,12 @@ class SuperwallPaywallActivity : AppCompatActivity() { ) } } else { - (viewStorageViewModel.retrieveView(ShimmerView.TAG) as ShimmerView) + cache.acquireShimmerView() } setupWith(shimmer, loading) } - viewStorageViewModel.storeView(key, this) + cache.storeView(key, this) } } @@ -286,7 +286,8 @@ class SuperwallPaywallActivity : AppCompatActivity() { } // Store the view again with the same key for this activity - viewStorageViewModel.storeView(key, currentPaywallView) + Superwall.instance.dependencyContainer.paywallManager.cache + .storeView(key, currentPaywallView) // Continue with normal activity setup using the restored view setupActivityWithView(currentPaywallView, presentationStyle) return diff --git a/superwall/src/main/java/com/superwall/sdk/paywall/view/webview/templating/TemplateLogic.kt b/superwall/src/main/java/com/superwall/sdk/paywall/view/webview/templating/TemplateLogic.kt index 1ec652ebd..9674b44fb 100644 --- a/superwall/src/main/java/com/superwall/sdk/paywall/view/webview/templating/TemplateLogic.kt +++ b/superwall/src/main/java/com/superwall/sdk/paywall/view/webview/templating/TemplateLogic.kt @@ -14,7 +14,6 @@ import com.superwall.sdk.paywall.view.webview.templating.models.ProductTemplate import kotlinx.serialization.json.Json object TemplateLogic { - suspend fun getBase64EncodedTemplates( json: Json, paywall: Paywall, 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 24332174d..db5aa250d 100644 --- a/superwall/src/main/java/com/superwall/sdk/web/WebPaywallRedeemer.kt +++ b/superwall/src/main/java/com/superwall/sdk/web/WebPaywallRedeemer.kt @@ -538,10 +538,12 @@ 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(), + ) } - factory.setWebEntitlements( - newEntitlements.filter { it.isActive }.toSet(), - ) // Trigger CustomerInfo merge customerInfoManager.updateMergedCustomerInfo() 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 index aabd96476..634d423cb 100644 --- a/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt @@ -5,15 +5,20 @@ 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.Assert.assertEquals @@ -21,6 +26,7 @@ import org.junit.Assert.assertNotNull import org.junit.Assert.assertNull import org.junit.Assert.assertSame import org.junit.Assert.assertTrue +import org.junit.After import org.junit.Before import org.junit.Test import org.junit.runner.RunWith @@ -39,8 +45,16 @@ class PaywallViewCacheTest { private fun keyOf(id: String) = PaywallCacheLogic.key(id, "en_US") + private lateinit var actorScope: CoroutineScope + private fun newCache(): PaywallViewCache = - PaywallViewCache(appCtx, storage, activityProvider, deviceHelper) + PaywallViewCache( + appCtx, + storage, + activityProvider, + deviceHelper, + actor = SequentialActor(PaywallCacheState(), actorScope), + ) @Before fun setup() { @@ -57,6 +71,12 @@ class PaywallViewCacheTest { object : ViewStorage { override val views = ConcurrentHashMap() } + actorScope = CoroutineScope(SupervisorJob() + Dispatchers.IO) + } + + @After + fun tearDown() { + actorScope.cancel() } // ------------------------------------------------------------------- @@ -389,8 +409,7 @@ class PaywallViewCacheTest { Then("the final read returns one of the assigned values") { val final = cache.activePaywallVcKey - assertNotNull(final) - assertTrue(final!!.startsWith("k_")) + assertTrue(final in (0 until 50).map { "k_$it" }) } } } @@ -412,10 +431,17 @@ class PaywallViewCacheTest { } (savers + removers).awaitAll() - Then("the cache does not crash and getAllPaywallViews is consistent") { - val views = cache.getAllPaywallViews() - // Final state depends on interleaving but must not throw - assertTrue(views.size in 0..ids.size) + 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, + ) } } } @@ -438,4 +464,101 @@ class PaywallViewCacheTest { } } } + + // ------------------------------------------------------------------- + // 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)) + } + } + } + } } diff --git a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt index 91c5a234c..b83eacca6 100644 --- a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt @@ -4,7 +4,6 @@ import com.superwall.sdk.And import com.superwall.sdk.Given import com.superwall.sdk.Then import com.superwall.sdk.When -import com.superwall.sdk.misc.primitives.StateActor import com.superwall.sdk.models.customer.CustomerInfo import com.superwall.sdk.models.entitlements.Entitlement import com.superwall.sdk.models.entitlements.SubscriptionStatus @@ -18,16 +17,12 @@ import com.superwall.sdk.store.abstractions.product.receipt.LatestSubscriptionSt import io.mockk.every import io.mockk.mockk import io.mockk.verify -import kotlinx.coroutines.Dispatchers -import kotlinx.coroutines.async -import kotlinx.coroutines.delay 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 -import kotlin.time.Duration.Companion.seconds /** * Comprehensive tests for the Entitlements class external API. @@ -82,16 +77,7 @@ class EntitlementsRefactorSafetyTest { Given("storage has no cached data") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) Then("status should be Unknown") { assertTrue(entitlements.status.value is SubscriptionStatus.Unknown) @@ -120,16 +106,7 @@ class EntitlementsRefactorSafetyTest { When("Entitlements is initialized") { val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) Then("corrupted status should be deleted from storage") { verify { storage.delete(StoredSubscriptionStatus) } @@ -154,16 +131,7 @@ class EntitlementsRefactorSafetyTest { When("Entitlements is initialized") { val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) Then("corrupted entitlements-by-product should be deleted") { verify { storage.delete(StoredEntitlementsByProductId) } @@ -181,16 +149,7 @@ class EntitlementsRefactorSafetyTest { Given("storage contains Inactive status") { val storage = mockStorage(storedStatus = SubscriptionStatus.Inactive) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) Then("status should be Inactive") { assertTrue(entitlements.status.value is SubscriptionStatus.Inactive) @@ -212,16 +171,7 @@ class EntitlementsRefactorSafetyTest { When("Entitlements is initialized") { val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) Then("entitlementsByProductId should contain the stored mappings") { assertEquals(productMap, entitlements.entitlementsByProductId) @@ -243,16 +193,7 @@ class EntitlementsRefactorSafetyTest { Given("a mix of active and inactive entitlements") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val activeE = Entitlement("active_one", isActive = true) val inactiveE = Entitlement("inactive_one", isActive = false) @@ -276,21 +217,12 @@ class EntitlementsRefactorSafetyTest { } @Test - fun `setSubscriptionStatus Active with all inactive entitlements becomes Inactive`() = + fun `setSubscriptionStatus Active with only inactive entitlements stays Active`() = runTest { Given("entitlements that are all inactive") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val inactiveE = Entitlement( id = "expired", @@ -319,16 +251,7 @@ class EntitlementsRefactorSafetyTest { Given("entitlements with Active status") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val e1 = Entitlement("first") val e2 = Entitlement("second") @@ -355,16 +278,7 @@ class EntitlementsRefactorSafetyTest { Given("entitlements cycling through states") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val e1 = Entitlement("premium") entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e1))) @@ -394,16 +308,7 @@ class EntitlementsRefactorSafetyTest { Given("entitlements in Active state with device entitlements") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("a")))) entitlements.activeDeviceEntitlements = setOf(Entitlement("device")) @@ -427,16 +332,7 @@ class EntitlementsRefactorSafetyTest { Given("entitlements in Unknown state") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("setting Inactive from Unknown") { entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) @@ -456,16 +352,7 @@ class EntitlementsRefactorSafetyTest { Given("entitlements subjected to rapid transitions") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val e1 = Entitlement("a") val e2 = Entitlement("b") val e3 = Entitlement("c") @@ -499,16 +386,7 @@ class EntitlementsRefactorSafetyTest { Given("entitlements with active device entitlements") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.activeDeviceEntitlements = setOf(Entitlement("device_premium")) When("setting Unknown status") { @@ -527,16 +405,7 @@ class EntitlementsRefactorSafetyTest { Given("entitlements with existing device entitlements") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.activeDeviceEntitlements = setOf(Entitlement("old")) When("setting new device entitlements") { @@ -557,16 +426,7 @@ class EntitlementsRefactorSafetyTest { Given("device entitlements set, then status goes Inactive") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.activeDeviceEntitlements = setOf(Entitlement("device")) When("setting Inactive") { @@ -594,16 +454,7 @@ class EntitlementsRefactorSafetyTest { redemptionResponse = webRedemption(webE), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("from_status")))) When("accessing all property") { @@ -633,16 +484,7 @@ class EntitlementsRefactorSafetyTest { ), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(activeE))) When("accessing inactive property") { @@ -652,11 +494,8 @@ class EntitlementsRefactorSafetyTest { assertTrue(inactive.any { it.id == "inactive_product" }) } And("it should not contain active entitlements") { - // active entitlement may appear in inactive if the exact object differs - // but we check the concept val activeIds = entitlements.active.map { it.id }.toSet() - val purelyInactive = inactive.filter { it.id !in activeIds } - assertTrue(purelyInactive.any { it.id == "inactive_product" }) + assertTrue(inactive.none { it.id in activeIds }) } } } @@ -668,16 +507,7 @@ class EntitlementsRefactorSafetyTest { Given("a fresh Entitlements with no data") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) Then("active should be empty") { assertTrue(entitlements.active.isEmpty()) @@ -695,16 +525,7 @@ class EntitlementsRefactorSafetyTest { Given("a fresh Entitlements instance") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val e1 = Entitlement("premium") val e2 = Entitlement("basic") val mapping = mapOf("prod_a" to setOf(e1), "prod_b" to setOf(e2)) @@ -733,16 +554,7 @@ class EntitlementsRefactorSafetyTest { Given("existing entitlements for a product") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val oldE = Entitlement("old") val newE = Entitlement("new") @@ -767,16 +579,7 @@ class EntitlementsRefactorSafetyTest { Given("a fresh Entitlements instance") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("adding an empty map") { entitlements.addEntitlementsByProductId(emptyMap()) @@ -798,16 +601,7 @@ class EntitlementsRefactorSafetyTest { Given("entitlements with product mappings") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.addEntitlementsByProductId(mapOf("prod1" to setOf(Entitlement("e1")))) When("taking a snapshot and then modifying the original") { @@ -844,16 +638,7 @@ class EntitlementsRefactorSafetyTest { ), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("querying with the exact full ID") { val result = entitlements.byProductId("sub:plan:offer") @@ -875,16 +660,7 @@ class EntitlementsRefactorSafetyTest { storedProductEntitlements = mapOf("monthly_sub" to setOf(e)), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("querying with a full ID that contains the subscription ID") { val result = entitlements.byProductId("monthly_sub:plan:offer") @@ -905,16 +681,7 @@ class EntitlementsRefactorSafetyTest { storedProductEntitlements = mapOf("known_product" to setOf(Entitlement("e"))), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("querying an unknown product") { val result = entitlements.byProductId("completely_unknown") @@ -940,16 +707,7 @@ class EntitlementsRefactorSafetyTest { ), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("querying product_a") { val result = entitlements.byProductId("product_a:plan") @@ -971,16 +729,7 @@ class EntitlementsRefactorSafetyTest { storedProductEntitlements = mapOf("com.app.product" to setOf(e)), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("querying the simple product ID") { val result = entitlements.byProductId("com.app.product") @@ -1011,16 +760,7 @@ class EntitlementsRefactorSafetyTest { ), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("querying multiple product IDs") { val result = entitlements.byProductIds(setOf("prod1", "prod2")) @@ -1043,16 +783,7 @@ class EntitlementsRefactorSafetyTest { storedProductEntitlements = mapOf("prod1" to setOf(Entitlement("e"))), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("querying with empty set") { val result = entitlements.byProductIds(emptySet()) @@ -1078,16 +809,7 @@ class EntitlementsRefactorSafetyTest { ), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("querying both products") { val result = entitlements.byProductIds(setOf("prod1", "prod2")) @@ -1110,16 +832,7 @@ class EntitlementsRefactorSafetyTest { storedProductEntitlements = mapOf("known_prod" to setOf(e1)), ) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("querying both") { val result = entitlements.byProductIds(setOf("known_prod", "unknown_prod")) @@ -1141,24 +854,12 @@ class EntitlementsRefactorSafetyTest { Given("Entitlements with backgroundScope for collector") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("setting Active status") { val activeE = setOf(Entitlement("persisted")) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(activeE)) - // Give collector time to process - async(Dispatchers.Default) { delay(1.seconds) }.await() - Then("storage write should have been called with the new status") { verify { storage.write( @@ -1177,20 +878,10 @@ class EntitlementsRefactorSafetyTest { Given("Entitlements with backgroundScope") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("setting Inactive status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) - async(Dispatchers.Default) { delay(1.seconds) }.await() Then("Inactive should be persisted") { verify { @@ -1224,16 +915,7 @@ class EntitlementsRefactorSafetyTest { ) val storage = mockStorage(redemptionResponse = redemption) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) Then("web should be empty") { assertTrue(entitlements.web.isEmpty()) @@ -1248,16 +930,7 @@ class EntitlementsRefactorSafetyTest { val webE = Entitlement("web_only", isActive = true, store = Store.STRIPE) val storage = mockStorage(redemptionResponse = webRedemption(webE)) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) Then("all should include web entitlements") { assertTrue(entitlements.all.contains(webE)) @@ -1275,16 +948,7 @@ class EntitlementsRefactorSafetyTest { val webE = Entitlement("web_sub", isActive = true, store = Store.STRIPE) val storage = mockStorage(redemptionResponse = webRedemption(webE)) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) Then("active should still contain web entitlements") { @@ -1303,16 +967,7 @@ class EntitlementsRefactorSafetyTest { Given("same entitlement ID from status and device sources") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val fromStatus = Entitlement("premium", isActive = true, store = Store.PLAY_STORE) val fromDevice = Entitlement("premium", isActive = true, store = Store.PLAY_STORE) @@ -1336,16 +991,7 @@ class EntitlementsRefactorSafetyTest { val webE = Entitlement("premium", isActive = true, store = Store.STRIPE) val storage = mockStorage(redemptionResponse = webRedemption(webE)) val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val statusE = Entitlement("premium", isActive = true, store = Store.PLAY_STORE) val deviceE = Entitlement("premium", isActive = true) @@ -1373,16 +1019,7 @@ class EntitlementsRefactorSafetyTest { Given("a richly-populated entitlement") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val now = Date() val future = Date(now.time + 86400000) val richE = @@ -1427,16 +1064,7 @@ class EntitlementsRefactorSafetyTest { Given("a single entitlement") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val e = Entitlement("solo") When("setting Active with single entitlement") { @@ -1456,16 +1084,7 @@ class EntitlementsRefactorSafetyTest { Given("100 entitlements") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val many = (1..100).map { Entitlement("e_$it") }.toSet() When("setting Active with all of them") { @@ -1487,16 +1106,7 @@ class EntitlementsRefactorSafetyTest { Given("Entitlements instance") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("setting status sequentially") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("a")))) @@ -1520,16 +1130,7 @@ class EntitlementsRefactorSafetyTest { Given("dynamically added product entitlements") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val e = Entitlement("dynamic") When("adding and then querying") { @@ -1551,16 +1152,7 @@ class EntitlementsRefactorSafetyTest { Given("only active entitlements") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val e = Entitlement("active", isActive = true) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(e))) @@ -1581,16 +1173,7 @@ class EntitlementsRefactorSafetyTest { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setWebEntitlements(setOf(webE1)) @@ -1612,16 +1195,7 @@ class EntitlementsRefactorSafetyTest { Given("entitlements with existing product mappings") { val storage = mockStorage() val entitlements = - StateActor( - createInitialEntitlementsState(storage), - backgroundScope, - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) val e1 = Entitlement("first") val e2 = Entitlement("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 69dabd574..3e2fc75d9 100644 --- a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsTest.kt @@ -4,7 +4,6 @@ import com.superwall.sdk.And import com.superwall.sdk.Given import com.superwall.sdk.Then import com.superwall.sdk.When -import com.superwall.sdk.misc.primitives.StateActor import com.superwall.sdk.models.customer.CustomerInfo import com.superwall.sdk.models.entitlements.Entitlement import com.superwall.sdk.models.entitlements.SubscriptionStatus @@ -52,29 +51,11 @@ class EntitlementsTest { ), ) entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("Entitlements is initialized") { val entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) Then("it should load the stored status") { assertEquals(storedStatus, entitlements.status.value) @@ -104,13 +85,7 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("setting active entitlement status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(activeEntitlements)) @@ -142,16 +117,7 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("setting active entitlement status with empty set") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(emptySet())) @@ -172,16 +138,7 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("setting NoActiveEntitlements status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Inactive) @@ -201,16 +158,7 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns null entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("setting Unknown status") { entitlements.setSubscriptionStatus(SubscriptionStatus.Unknown) @@ -242,16 +190,7 @@ class EntitlementsTest { ) every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("creating a new Entitlements instance") { Then("it should return correct entitlements for each product") { @@ -286,16 +225,7 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("querying with subscription_monthly colon p1m colon freetrial") { val result = entitlements.byProductId("subscription_monthly:p1m:freetrial") @@ -323,16 +253,7 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("setting active device entitlements to only the active one") { entitlements.activeDeviceEntitlements = setOf(activeEntitlement) @@ -371,16 +292,7 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) When("no active device entitlements are set") { // activeDeviceEntitlements not set, should be empty @@ -413,13 +325,7 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.activeDeviceEntitlements = setOf(activeEntitlement) When("subscription status is set to Inactive") { @@ -456,13 +362,7 @@ class EntitlementsTest { every { storage.read(StoredSubscriptionStatus) } returns null every { storage.read(StoredEntitlementsByProductId) } returns productEntitlements entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("setting both status and device entitlements") { entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusActiveEntitlement))) @@ -519,16 +419,7 @@ class EntitlementsTest { When("accessing the web property") { entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) Then("it should return only active web entitlements") { assertEquals(setOf(webEntitlement1, webEntitlement2), entitlements.web) @@ -564,16 +455,7 @@ class EntitlementsTest { When("accessing the web property") { entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) Then("it should return only active web entitlements") { assertEquals(setOf(activeWebEntitlement), entitlements.web) @@ -593,16 +475,7 @@ class EntitlementsTest { When("accessing the web property") { entitlements = - StateActor( - createInitialEntitlementsState(storage), - CoroutineScope(Dispatchers.Default), - ).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = actor.scope, - ) - } + makeEntitlements(storage, CoroutineScope(Dispatchers.Default)) Then("it should return empty set") { assertTrue(entitlements.web.isEmpty()) @@ -643,13 +516,7 @@ class EntitlementsTest { When("setting subscription status (simulating external PC)") { entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusEntitlement))) Then("active should contain both status and web entitlements") { @@ -695,13 +562,7 @@ class EntitlementsTest { When("external PC sets status with only its entitlements") { entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) // External PC sets status (like RC does) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(rcEntitlement))) @@ -745,13 +606,7 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) When("external PC reads web entitlements and merges them into status") { // This simulates what the updated RC controller does: @@ -809,13 +664,7 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(playEntitlement))) When("status is reset to Inactive (simulating sign out)") { @@ -862,13 +711,7 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) // Initial state entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("old_play")))) @@ -919,13 +762,7 @@ class EntitlementsTest { every { storage.read(StoredEntitlementsByProductId) } returns null entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(Entitlement("userA_play")))) When("user B identifies and web entitlements are updated") { @@ -986,13 +823,7 @@ class EntitlementsTest { When("all three sources have different entitlements") { entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusEntitlement))) entitlements.activeDeviceEntitlements = setOf(deviceEntitlement) @@ -1036,13 +867,7 @@ class EntitlementsTest { When("both sources have entitlement with same ID") { entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + makeEntitlements(storage, backgroundScope) entitlements.setSubscriptionStatus(SubscriptionStatus.Active(setOf(statusPremium))) Then("active should deduplicate and contain only one premium entitlement") { @@ -1080,13 +905,7 @@ class EntitlementsTest { When("status is set to Unknown") { entitlements = - StateActor(createInitialEntitlementsState(storage), backgroundScope).let { actor -> - Entitlements( - storage = storage, - actor = actor, - actorScope = backgroundScope, - ) - } + 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 632efb438..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 @@ -1212,13 +1212,12 @@ class TestModeTest { // Subscription status reflects the active selection on both // TestMode itself and the entitlements collaborator. - val expectedStatus = manager.buildSubscriptionStatus() - assertTrue( - "Expected SubscriptionStatus.Active", - expectedStatus is SubscriptionStatus.Active, - ) - assertEquals(expectedStatus, manager.overriddenSubscriptionStatus) - verify(exactly = 1) { entitlements.setSubscriptionStatus(expectedStatus) } + 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. 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 36e93d35d..e8e01f46e 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,7 @@ 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.Before import org.junit.Test @@ -131,6 +133,7 @@ class WebPaywallRedeemerTest { ) }, var getIntegrationPropsFn: () -> Map = { emptyMap() }, + var setWebEntitlementsFn: (Set) -> Unit = {}, ) : WebPaywallRedeemer.Factory { override fun willRedeemLink() = willRedeemLinkFn() @@ -150,7 +153,7 @@ class WebPaywallRedeemerTest { override fun internallySetSubscriptionStatus(status: SubscriptionStatus) = this@WebPaywallRedeemerTest.setSubscriptionStatus(status) - override fun setWebEntitlements(entitlements: Set) {} + override fun setWebEntitlements(entitlements: Set) = setWebEntitlementsFn(entitlements) override suspend fun isPaywallVisible(): Boolean = this@WebPaywallRedeemerTest.isPaywallVisible() @@ -241,6 +244,8 @@ class WebPaywallRedeemerTest { ) } returns Either.Success(response) + val published = java.util.concurrent.CopyOnWriteArrayList>() + When("creating redeemer and advancing scheduler") { redeemer = WebPaywallRedeemer( @@ -250,7 +255,7 @@ class WebPaywallRedeemerTest { network, storage, customerInfoManager = mockk(relaxed = true), - factory = TestFactory(), + factory = TestFactory(setWebEntitlementsFn = { published.add(it) }), ) testScheduler.advanceUntilIdle() @@ -258,9 +263,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()) + } } } } From 6b5b2252572996e2fb5ba789b59f4f0532aa7c52 Mon Sep 17 00:00:00 2001 From: Ian Rumac Date: Fri, 25 Sep 2026 14:42:33 +0200 Subject: [PATCH 04/10] Route activity view access through PaywallViewRegistry SuperwallPaywallActivity and DebugView now reach paywall views through an internal PaywallViewRegistry obtained from DependencyContainer.makeViewRegistry(), instead of reaching into paywallManager.cache or ViewStorage directly. Writes go through the cache so it and ViewStorage stay in sync; reads use ViewStorage, which survives Activity recreation. PaywallManager's cache is private again. Adds an on-device test proving the loading and shimmer views can be built on a thread without a Looper, as the cache actor does, and still draw and animate on main. Co-Authored-By: Claude Opus 5.5 --- .../view/OffMainViewConstructionTest.kt | 102 ++++++++++++++++++ .../java/com/superwall/sdk/debug/DebugView.kt | 5 +- .../sdk/dependencies/DependencyContainer.kt | 7 ++ .../sdk/paywall/manager/PaywallManager.kt | 9 +- .../paywall/manager/PaywallViewRegistry.kt | 46 ++++++++ .../paywall/view/SuperwallPaywallActivity.kt | 22 ++-- .../paywall/manager/PaywallViewCacheTest.kt | 23 ++++ 7 files changed, 196 insertions(+), 18 deletions(-) create mode 100644 superwall/src/androidTest/java/com/superwall/sdk/paywall/view/OffMainViewConstructionTest.kt create mode 100644 superwall/src/main/java/com/superwall/sdk/paywall/manager/PaywallViewRegistry.kt 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..cf6663218 --- /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 cacheAcquireFromMainBuildsViewsOnTheActorThread() { + 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/debug/DebugView.kt b/superwall/src/main/java/com/superwall/sdk/debug/DebugView.kt index 9ca50698a..2161ab9ac 100644 --- a/superwall/src/main/java/com/superwall/sdk/debug/DebugView.kt +++ b/superwall/src/main/java/com/superwall/sdk/debug/DebugView.kt @@ -923,7 +923,8 @@ internal class DebugViewActivity : AppCompatActivity() { view: View, ) { val key = UUID.randomUUID().toString() - Superwall.instance.dependencyContainer.paywallManager.cache + Superwall.instance.dependencyContainer + .makeViewRegistry() .storeView(key, view) val intent = @@ -961,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 3c2654faf..cd4eafc58 100644 --- a/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt +++ b/superwall/src/main/java/com/superwall/sdk/dependencies/DependencyContainer.kt @@ -79,6 +79,7 @@ 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 @@ -1124,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 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 d75e20fb7..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 @@ -30,11 +30,7 @@ class PaywallManager( private var _cache: PaywallViewCache? = null - /** - * The single cache instance. Exposed so Activities that write to - * [com.superwall.sdk.paywall.view.ViewStorage] can go through it instead. - */ - internal val cache: PaywallViewCache + private val cache: PaywallViewCache get() { if (_cache == null) { _cache = createCache() @@ -42,6 +38,9 @@ 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 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/view/SuperwallPaywallActivity.kt b/superwall/src/main/java/com/superwall/sdk/paywall/view/SuperwallPaywallActivity.kt index f7601db0b..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 @@ -157,7 +157,8 @@ class SuperwallPaywallActivity : AppCompatActivity() { } return launchPaywallActivity(context, intent).onFailure { - Superwall.instance.dependencyContainer.paywallManager.cache + Superwall.instance.dependencyContainer + .makeViewRegistry() .removeView(key) view.clearActivityLaunchState() } @@ -166,13 +167,13 @@ class SuperwallPaywallActivity : AppCompatActivity() { private fun PaywallView.prepareViewForDisplay(key: String) { webView.enableBackgroundRendering() webView.attach(this) - val cache = Superwall.instance.dependencyContainer.paywallManager.cache + 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. 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 = cache.acquireLoadingView() + 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 { - cache.acquireShimmerView() + registry.acquireShimmerView() } setupWith(shimmer, loading) } - cache.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,8 +287,7 @@ class SuperwallPaywallActivity : AppCompatActivity() { } // Store the view again with the same key for this activity - Superwall.instance.dependencyContainer.paywallManager.cache - .storeView(key, currentPaywallView) + viewRegistry.storeView(key, currentPaywallView) // Continue with normal activity setup using the restored view setupActivityWithView(currentPaywallView, presentationStyle) return @@ -845,7 +845,7 @@ class SuperwallPaywallActivity : AppCompatActivity() { if (pv != null) { ( Superwall.instance.dependencyContainer - .makeViewStore() + .makeViewRegistry() .retrieveView(pv) as? PaywallView? )?.cleanup() } 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 index 634d423cb..e8f81689b 100644 --- a/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt @@ -561,4 +561,27 @@ class PaywallViewCacheTest { } } } + + @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")) + } + } + } + } } From 7c8446f883a3c41acc4fc010b43dc2a9d5a6d3b7 Mon Sep 17 00:00:00 2001 From: Ian Rumac Date: Fri, 25 Sep 2026 14:42:33 +0200 Subject: [PATCH 05/10] Keep status entitlements in all after a cold start createInitialEntitlementsState replayed the saved status before the stored product entitlements. AddProductEntitlements replaces allTracked, so status entitlements not tied to a product dropped out of `all` while still in `active`, until the next status update. The old startup code never replaced that set. Restore product entitlements first. Co-Authored-By: Claude Opus 5.5 --- .../superwall/sdk/store/EntitlementsState.kt | 12 +++++--- .../store/EntitlementsRefactorSafetyTest.kt | 29 +++++++++++++++++++ 2 files changed, 37 insertions(+), 4 deletions(-) diff --git a/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt index ed1f94450..6785a1aff 100644 --- a/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt +++ b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt @@ -159,6 +159,14 @@ internal fun createInitialEntitlementsState(storage: Storage): EntitlementsState var state = EntitlementsState() + // Restore product entitlements BEFORE the status. AddProductEntitlements + // replaces allTracked, so replaying it after SetActive would drop status + // entitlements that are not tied to a product from `all` until the next + // status update. The old startup code never replaced that set. + if (productEntitlements != null) { + state = EntitlementsState.Updates.AddProductEntitlements(productEntitlements).reduce(state) + } + // Replay status to populate backingActive/allTracked correctly if (status != null) { state = @@ -175,10 +183,6 @@ internal fun createInitialEntitlementsState(storage: Storage): EntitlementsState } } - if (productEntitlements != null) { - state = EntitlementsState.Updates.AddProductEntitlements(productEntitlements).reduce(state) - } - // Restore web entitlements from latest redemption response val webEntitlements = try { diff --git a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt index b83eacca6..022758f4e 100644 --- a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt @@ -118,6 +118,35 @@ class EntitlementsRefactorSafetyTest { } } + @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 `init with corrupted StoredEntitlementsByProductId deletes and continues`() = runTest { From c10cfac7e00144d15509bfa2a27da7f171a076a3 Mon Sep 17 00:00:00 2001 From: Ian Rumac Date: Wed, 30 Sep 2026 13:50:27 +0200 Subject: [PATCH 06/10] Minor fixes --- .../paywall/manager/PaywallViewCacheTest.kt | 2 +- .../sdk/web/WebPaywallRedeemerTest.kt | 82 +++++++++++++++++++ 2 files changed, 83 insertions(+), 1 deletion(-) 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 index e8f81689b..22e499404 100644 --- a/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/paywall/manager/PaywallViewCacheTest.kt @@ -21,12 +21,12 @@ 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.After import org.junit.Before import org.junit.Test import org.junit.runner.RunWith 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 e8e01f46e..41460a265 100644 --- a/superwall/src/test/java/com/superwall/sdk/web/WebPaywallRedeemerTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/web/WebPaywallRedeemerTest.kt @@ -44,6 +44,7 @@ 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 @@ -938,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" }) + } + } + } + } } From ff76d54a2526516c2da531545c85e1d357808968 Mon Sep 17 00:00:00 2001 From: Ian Rumac Date: Wed, 30 Sep 2026 14:16:43 +0200 Subject: [PATCH 07/10] Merge product entitlements into allTracked instead of replacing AddProductEntitlements overwrote allTracked, so a status-only entitlement dropped out of `all` once product entitlements loaded while still showing in `active`. Also rename the off-main cache acquire test to match what it asserts. Co-Authored-By: Claude Opus 5.5 --- .../view/OffMainViewConstructionTest.kt | 2 +- .../superwall/sdk/store/EntitlementsState.kt | 8 +++--- .../store/EntitlementsRefactorSafetyTest.kt | 26 +++++++++++++++++++ 3 files changed, 30 insertions(+), 6 deletions(-) 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 index cf6663218..657ffe9a4 100644 --- a/superwall/src/androidTest/java/com/superwall/sdk/paywall/view/OffMainViewConstructionTest.kt +++ b/superwall/src/androidTest/java/com/superwall/sdk/paywall/view/OffMainViewConstructionTest.kt @@ -75,7 +75,7 @@ class OffMainViewConstructionTest { } @Test - fun cacheAcquireFromMainBuildsViewsOnTheActorThread() { + fun cacheAcquireFromMainReturnsUsableViewsWithoutDeadlocking() { val scope = CoroutineScope(SupervisorJob() + Dispatchers.IO) try { val cache = diff --git a/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt index 6785a1aff..860ab5c21 100644 --- a/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt +++ b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt @@ -111,7 +111,7 @@ data class EntitlementsState( idToEntitlements.mapValues { (_, v) -> v.toSet() } state.copy( entitlementsByProduct = newProducts, - allTracked = newProducts.values.flatten().toSet(), + allTracked = state.allTracked + newProducts.values.flatten(), ) }) @@ -159,10 +159,8 @@ internal fun createInitialEntitlementsState(storage: Storage): EntitlementsState var state = EntitlementsState() - // Restore product entitlements BEFORE the status. AddProductEntitlements - // replaces allTracked, so replaying it after SetActive would drop status - // entitlements that are not tied to a product from `all` until the next - // status update. The old startup code never replaced that set. + // Restore product entitlements, then the status. Both merge into + // allTracked, so status entitlements not tied to a product stay in `all`. if (productEntitlements != null) { state = EntitlementsState.Updates.AddProductEntitlements(productEntitlements).reduce(state) } diff --git a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt index 022758f4e..79d6180ef 100644 --- a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt @@ -147,6 +147,32 @@ class EntitlementsRefactorSafetyTest { } } + @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 { From 2d272fb9fea482e991dccef45ad77aff3b26e72d Mon Sep 17 00:00:00 2001 From: Ian Rumac Date: Wed, 30 Sep 2026 16:34:24 +0200 Subject: [PATCH 08/10] Fix delegate status 'from', keep allTracked status-only, pin delegate override - Take the delegate's previous status from the status flow. Entitlements now persists before the listener runs, so reading storage gave from == to. - AddProductEntitlements no longer writes allTracked; `all` already unions product entitlements, and this avoids stale entries on a remapped product. - Add SubscriptionStatusDelegateOverrideTest for setting the status from inside subscriptionStatusDidChange (add, remove, replace, grant). - Stub the suspend PaywallViewCache.save with coEvery in PaywallManagerExperimentIsolationTest so unit tests compile. Co-Authored-By: Claude Opus 5.5 --- .../main/java/com/superwall/sdk/Superwall.kt | 14 +- .../superwall/sdk/store/EntitlementsState.kt | 11 +- .../SubscriptionStatusDelegateOverrideTest.kt | 313 ++++++++++++++++++ .../PaywallManagerExperimentIsolationTest.kt | 2 +- .../store/EntitlementsRefactorSafetyTest.kt | 63 +++- 5 files changed, 388 insertions(+), 15 deletions(-) create mode 100644 superwall/src/test/java/com/superwall/sdk/SubscriptionStatusDelegateOverrideTest.kt diff --git a/superwall/src/main/java/com/superwall/sdk/Superwall.kt b/superwall/src/main/java/com/superwall/sdk/Superwall.kt index b8b76fd7a..bcec483df 100644 --- a/superwall/src/main/java/com/superwall/sdk/Superwall.kt +++ b/superwall/src/main/java/com/superwall/sdk/Superwall.kt @@ -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/store/EntitlementsState.kt b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt index 860ab5c21..7fd9fa8bb 100644 --- a/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt +++ b/superwall/src/main/java/com/superwall/sdk/store/EntitlementsState.kt @@ -18,7 +18,7 @@ data class EntitlementsState( val backingActive: Set = emptySet(), /** Active web entitlements from the latest redemption response. */ val webEntitlements: Set = emptySet(), - /** Tracks all entitlements seen from status updates + product updates. */ + /** Tracks all entitlements seen from status updates. */ val allTracked: Set = emptySet(), ) { // -- Derived properties -- @@ -109,10 +109,7 @@ data class EntitlementsState( val newProducts = state.entitlementsByProduct + idToEntitlements.mapValues { (_, v) -> v.toSet() } - state.copy( - entitlementsByProduct = newProducts, - allTracked = state.allTracked + newProducts.values.flatten(), - ) + state.copy(entitlementsByProduct = newProducts) }) data class SetDeviceEntitlements( @@ -159,8 +156,8 @@ internal fun createInitialEntitlementsState(storage: Storage): EntitlementsState var state = EntitlementsState() - // Restore product entitlements, then the status. Both merge into - // allTracked, so status entitlements not tied to a product stay in `all`. + // 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) } 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..8a8036688 --- /dev/null +++ b/superwall/src/test/java/com/superwall/sdk/SubscriptionStatusDelegateOverrideTest.kt @@ -0,0 +1,313 @@ +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") { + verify { + storage.write( + StoredSubscriptionStatus, + match { it.ids() == setOf("pro", "custom") }, + ) + } + } + } + } + } + + @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") { + verify { + storage.write(StoredSubscriptionStatus, match { it.ids() == setOf("pro") }) + } + } + // 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/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/store/EntitlementsRefactorSafetyTest.kt b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt index 79d6180ef..2663ab8ec 100644 --- a/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/store/EntitlementsRefactorSafetyTest.kt @@ -17,6 +17,9 @@ import com.superwall.sdk.store.abstractions.product.receipt.LatestSubscriptionSt 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 @@ -1245,7 +1248,65 @@ class EntitlementsRefactorSafetyTest { } @Test - fun `addEntitlementsByProductId clears and rebuilds _all`() = + 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() From 4933348082fd29e28a9ed8c72e9f054095c4dab7 Mon Sep 17 00:00:00 2001 From: Ian Rumac Date: Fri, 2 Oct 2026 11:55:01 +0200 Subject: [PATCH 09/10] Assert the last persisted status in delegate override tests Co-Authored-By: Claude Opus 5.5 --- .../sdk/SubscriptionStatusDelegateOverrideTest.kt | 15 ++++++--------- 1 file changed, 6 insertions(+), 9 deletions(-) diff --git a/superwall/src/test/java/com/superwall/sdk/SubscriptionStatusDelegateOverrideTest.kt b/superwall/src/test/java/com/superwall/sdk/SubscriptionStatusDelegateOverrideTest.kt index 8a8036688..6afe1655f 100644 --- a/superwall/src/test/java/com/superwall/sdk/SubscriptionStatusDelegateOverrideTest.kt +++ b/superwall/src/test/java/com/superwall/sdk/SubscriptionStatusDelegateOverrideTest.kt @@ -147,12 +147,9 @@ class SubscriptionStatusDelegateOverrideTest { ) } And("the overridden status is what gets persisted last") { - verify { - storage.write( - StoredSubscriptionStatus, - match { it.ids() == setOf("pro", "custom") }, - ) - } + val writes = mutableListOf() + verify { storage.write(StoredSubscriptionStatus, capture(writes)) } + assertEquals(setOf("pro", "custom"), writes.last().ids()) } } } @@ -201,9 +198,9 @@ class SubscriptionStatusDelegateOverrideTest { ) } And("the narrowed status is what gets persisted last") { - verify { - storage.write(StoredSubscriptionStatus, match { it.ids() == setOf("pro") }) - } + 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 From b8eea44e3dafb43fdb454592ea79536c4be04fff Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" Date: Fri, 2 Oct 2026 10:13:49 +0000 Subject: [PATCH 10/10] Update coverage badge [skip ci] --- .github/badges/branches.svg | 2 +- .github/badges/jacoco.svg | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) 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