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