package com.suno.android.common_data.billing import app.cash.turbine.test import com.suno.android.common_core_utils.environment.UserPrefsDataStoreManager import com.suno.android.common_data.billing.SelectedModelProvider.Companion.FALLBACK_MODEL_KEY import com.suno.android.common_data.billing.SelectedModelProvider.Companion.FALLBACK_MODEL_NAME import com.suno.android.common_networking.remote.entities.SubscriptionInfoResponse import com.suno.android.common_networking.remote.session.Model import io.mockk.every import io.mockk.mockk import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.test.runTest import org.junit.Assert.assertEquals import org.junit.Assert.assertNotNull import org.junit.Test class SelectedModelProviderImplTest { private val billingStateFlow = MutableStateFlow(null) private val billingRepo = mockk(relaxed = true) { every { billingStateFlow() } returns billingStateFlow } private val selectedModelNameFlow = MutableStateFlow(null) private val userPrefsDataStoreManager = mockk(relaxed = true) { every { getSelectedModel() } returns selectedModelNameFlow } private val subject = SelectedModelProviderImpl( billingRepo = billingRepo, userPrefsDataStoreManager = userPrefsDataStoreManager, ) @Test fun `given null billing info when getSelectedModelFlow then returns fallback model`() = runTest { billingStateFlow.value = null selectedModelNameFlow.value = "v5" subject.getSelectedModelFlow().test { val model = awaitItem() assertEquals(FALLBACK_MODEL_KEY, model.externalKey) assertEquals(FALLBACK_MODEL_NAME, model.name) assertEquals(true, model.isDefaultModel) } } @Test fun `given empty models list when getSelectedModelFlow then returns fallback model`() = runTest { val billingInfo = createBillingInfo(models = emptyList()) billingStateFlow.value = billingInfo selectedModelNameFlow.value = "v5" subject.getSelectedModelFlow().test { val model = awaitItem() assertEquals(FALLBACK_MODEL_KEY, model.externalKey) assertEquals(FALLBACK_MODEL_NAME, model.name) } } @Test fun `given matching model in billing info when getSelectedModelFlow then returns that model`() = runTest { val testModel = createModel(name = "v5", externalKey = "chirp-crow") val billingInfo = createBillingInfo(models = listOf(testModel)) billingStateFlow.value = billingInfo selectedModelNameFlow.value = "v5" subject.getSelectedModelFlow().test { val model = awaitItem() assertEquals("chirp-crow", model.externalKey) assertEquals("v5", model.name) } } @Test fun `given model with null externalKey when getSelectedModelFlow then returns fallback model`() = runTest { val testModel = createModel(name = "v5", externalKey = null) val billingInfo = createBillingInfo(models = listOf(testModel)) billingStateFlow.value = billingInfo selectedModelNameFlow.value = "v5" subject.getSelectedModelFlow().test { val model = awaitItem() assertEquals(FALLBACK_MODEL_KEY, model.externalKey) } } @Test fun `given selected model name not in billing info when getSelectedModelFlow then returns fallback model`() = runTest { val testModel = createModel(name = "v4", externalKey = "chirp-v4") val billingInfo = createBillingInfo(models = listOf(testModel)) billingStateFlow.value = billingInfo selectedModelNameFlow.value = "v5" subject.getSelectedModelFlow().test { val model = awaitItem() assertEquals(FALLBACK_MODEL_KEY, model.externalKey) } } @Test fun `given null selected model name when getSelectedModelFlow then returns fallback model`() = runTest { val testModel = createModel(name = "v5", externalKey = "chirp-crow") val billingInfo = createBillingInfo(models = listOf(testModel)) billingStateFlow.value = billingInfo selectedModelNameFlow.value = null subject.getSelectedModelFlow().test { val model = awaitItem() assertEquals(FALLBACK_MODEL_KEY, model.externalKey) } } @Test fun `given multiple models when getSelectedModelFlow then returns correct model by name`() = runTest { val model1 = createModel(name = "v3.5", externalKey = "chirp-v3-5") val model2 = createModel(name = "v4", externalKey = "chirp-v4") val model3 = createModel(name = "v5", externalKey = "chirp-crow") val billingInfo = createBillingInfo(models = listOf(model1, model2, model3)) billingStateFlow.value = billingInfo selectedModelNameFlow.value = "v4" subject.getSelectedModelFlow().test { val model = awaitItem() assertEquals("chirp-v4", model.externalKey) assertEquals("v4", model.name) } } @Test fun `given model selection changes when getSelectedModelFlow then emits new model`() = runTest { val model1 = createModel(name = "v3.5", externalKey = "chirp-v3-5") val model2 = createModel(name = "v5", externalKey = "chirp-crow") val billingInfo = createBillingInfo(models = listOf(model1, model2)) billingStateFlow.value = billingInfo selectedModelNameFlow.value = "v3.5" subject.getSelectedModelFlow().test { val firstModel = awaitItem() assertEquals("chirp-v3-5", firstModel.externalKey) selectedModelNameFlow.value = "v5" val secondModel = awaitItem() assertEquals("chirp-crow", secondModel.externalKey) } } @Test fun `given billing info changes when getSelectedModelFlow then emits updated model`() = runTest { val model1 = createModel(name = "v5", externalKey = "chirp-crow-v1") val billingInfo1 = createBillingInfo(models = listOf(model1)) billingStateFlow.value = billingInfo1 selectedModelNameFlow.value = "v5" subject.getSelectedModelFlow().test { val firstModel = awaitItem() assertEquals("chirp-crow-v1", firstModel.externalKey) val model2 = createModel(name = "v5", externalKey = "chirp-crow-v2") val billingInfo2 = createBillingInfo(models = listOf(model2)) billingStateFlow.value = billingInfo2 val secondModel = awaitItem() assertEquals("chirp-crow-v2", secondModel.externalKey) } } @Test fun `given billing info changes to null when getSelectedModelFlow then emits fallback model`() = runTest { val model1 = createModel(name = "v5", externalKey = "chirp-crow") val billingInfo = createBillingInfo(models = listOf(model1)) billingStateFlow.value = billingInfo selectedModelNameFlow.value = "v5" subject.getSelectedModelFlow().test { val firstModel = awaitItem() assertEquals("chirp-crow", firstModel.externalKey) billingStateFlow.value = null val secondModel = awaitItem() assertEquals(FALLBACK_MODEL_KEY, secondModel.externalKey) } } @Test fun `given fallback model when checking properties then has expected values`() = runTest { billingStateFlow.value = null selectedModelNameFlow.value = "v5" subject.getSelectedModelFlow().test { val model = awaitItem() assertEquals(FALLBACK_MODEL_KEY, model.externalKey) assertEquals(FALLBACK_MODEL_NAME, model.name) assertEquals(5, model.majorVersion) assertEquals(true, model.isDefaultModel) assertEquals(true, model.canUse) assertNotNull(model.maxLengths) assertEquals(500, model.maxLengths?.gptDescriptionPrompt) assertEquals(1000, model.maxLengths?.negativeTags) assertEquals(5000, model.maxLengths?.prompt) assertEquals(1000, model.maxLengths?.tags) assertEquals(100, model.maxLengths?.title) } } @Test fun `given model selection changes to unavailable model when getSelectedModelFlow then emits fallback`() = runTest { val model1 = createModel(name = "v5", externalKey = "chirp-crow") val billingInfo = createBillingInfo(models = listOf(model1)) billingStateFlow.value = billingInfo selectedModelNameFlow.value = "v5" subject.getSelectedModelFlow().test { val firstModel = awaitItem() assertEquals("chirp-crow", firstModel.externalKey) selectedModelNameFlow.value = "v6" val secondModel = awaitItem() assertEquals(FALLBACK_MODEL_KEY, secondModel.externalKey) } } private fun createModel( name: String, externalKey: String?, isDefaultModel: Boolean = false, ): Model = Model( canUse = true, capabilities = emptyList(), features = emptyList(), description = "Test model", externalKey = externalKey, id = "model-id", majorVersion = 3, maxLengths = null, name = name, isDefaultModel = isDefaultModel, badges = emptyList(), ) private fun createBillingInfo( models: List = emptyList(), ): SubscriptionInfoResponse = mockk(relaxed = true) { every { this@mockk.models } returns models } }