package com.suno.android.ui.screens.orpheus import app.cash.turbine.test import arrow.core.Either import arrow.retrofit.adapter.either.networkhandling.CallError import com.suno.android.clip.UpdateClipReactionUseCase import com.suno.android.common_analytics.listening_source.ListeningSourceCache import com.suno.android.common_analytics.managers.AnalyticsManager import com.suno.android.common_analytics.orpheus.OrpheusAnalyticsManager import com.suno.android.common_core_utils.Id import com.suno.android.common_core_utils.asId import com.suno.android.common_core_utils.constants.ReactionType import com.suno.android.common_core_utils.constants.SunoMediaType import com.suno.android.common_core_utils.environment.ThemeMode import com.suno.android.common_core_utils.environment.UserPrefsDataStoreManager import com.suno.android.common_core_utils.model.AsyncData import com.suno.android.common_core_utils.model.UiString import com.suno.android.common_core_utils.model.Url import com.suno.android.common_core_utils.model.UserHandle import com.suno.android.common_data.billing.SelectedModelProvider import com.suno.android.common_data.billing.SunoBillingRepo import com.suno.android.common_data.generation.SongGenerationStoreImpl import com.suno.android.common_data.mappers.clips.ClipStatus import com.suno.android.common_data.mappers.clips.LocalClipData import com.suno.android.common_data.mappers.clips.SongListData import com.suno.android.common_data.mappers.projects.BaseProjectMetadata import com.suno.android.common_data.mappers.projects.ProjectData import com.suno.android.common_data.mappers.projects.ProjectMetadata import com.suno.android.common_data.repos.GenerationRepository import com.suno.android.common_data.repos.ShareLinkRepository import com.suno.android.common_data.repos.orpheus.OrpheusChatRepository import com.suno.android.common_data.repos.orpheus.OrpheusSessionError import com.suno.android.common_data.repos.orpheus.OrpheusSessionStore import com.suno.android.common_data.repos.orpheus.models.MessageContentType import com.suno.android.common_data.repos.orpheus.models.MessageRole import com.suno.android.common_data.repos.orpheus.models.OrpheusMessage import com.suno.android.common_data.repos.orpheus.models.OrpheusMessage.MessageStatus.AssistantMessageStatus import com.suno.android.common_data.repos.orpheus.models.OrpheusSession import com.suno.android.common_mvi.MviProcessorFactory import com.suno.android.common_networking.remote.entities.SubscriptionInfoResponse import com.suno.android.common_networking.remote.session.Model import com.suno.android.common_ui.components.bottom_sheet.SharePlatformConstants import com.suno.android.common_ui.components.omni.RepeatMode import com.suno.android.gating.FeatureManager import com.suno.android.gating.statsig.StatsigFeatureDataSource import com.suno.android.media.MediaManager import com.suno.android.media.MediaPlayerState import com.suno.android.ui.screens.orpheus.components.projects.OrpheusProjectsController import com.suno.android.usecase.GetStringFromResourcesUseCase import io.mockk.coEvery import io.mockk.coVerify import io.mockk.every import io.mockk.mockk import io.mockk.verify import kotlinx.collections.immutable.persistentListOf import kotlinx.collections.immutable.toImmutableMap import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.test.StandardTestDispatcher import kotlinx.coroutines.test.advanceUntilIdle import kotlinx.coroutines.test.resetMain import kotlinx.coroutines.test.runTest import kotlinx.coroutines.test.setMain import org.junit.After import org.junit.Assert.assertEquals import org.junit.Assert.assertNotNull import org.junit.Assert.assertNull import org.junit.Assert.assertTrue import org.junit.Before import org.junit.Test import java.util.Date import kotlin.time.Duration import kotlin.time.Duration.Companion.seconds class OrpheusChatVMTest { private val processorFactory = mockk() private val orpheusChatRepository = mockk(relaxed = true) private val sessionStore = mockk(relaxed = true) private val listeningSourceCache = mockk(relaxed = true) private val mediaManager = mockk(relaxed = true) private val generationRepository = mockk(relaxed = true) private val analyticsManager = mockk(relaxed = true) private val statsigManager = mockk(relaxed = true) private val updateClipReactionUseCase = mockk(relaxed = true) private val billingRepo = mockk(relaxed = true) private val selectedModelProvider = mockk(relaxed = true) private val shareLinkRepository = mockk(relaxed = true) private val projectControllerFactory = mockk(relaxed = true) private val featureManager = mockk(relaxed = true) private val orpheusAnalyticsManager = mockk(relaxed = true) private val getStringFromResourcesUseCase = mockk(relaxed = true) private val userPrefsDataStoreManager = mockk(relaxed = true) private val testMessageId = Id("test-msg-1") private val messagesFlow = MutableStateFlow(AsyncData.Ready(persistentListOf())) private val mediaPlayerStateFlow = MutableStateFlow(createMediaPlayerState()) private val songGenerationStateFlow = MutableStateFlow(SongGenerationStoreImpl.SongGenerationState()) private val billingStateFlow = MutableStateFlow(null) private val selectedModelFlow = MutableStateFlow( Model( canUse = true, capabilities = emptyList(), features = emptyList(), description = "Test model", externalKey = "chirp-crow", id = "test-model-id", majorVersion = 4, maxLengths = null, name = "chirp-crow", isDefaultModel = false, badges = emptyList(), ), ) private val testSession = OrpheusSession( sessionId = Id("test-session"), createdAt = 0.seconds, project = object : BaseProjectMetadata { override val id = "test-project".asId() override val name = "Test Project" }, ) private val sessionFlow = MutableStateFlow>(AsyncData.Ready(testSession)) private val themeModeFlow = MutableStateFlow(ThemeMode.SYSTEM) private lateinit var subject: OrpheusChatVM @Before fun setUp() { Dispatchers.setMain(StandardTestDispatcher()) every { orpheusChatRepository.messages } returns messagesFlow every { mediaManager.mediaPlayerFlow() } returns mediaPlayerStateFlow every { generationRepository.songGenerationStateFlow() } returns songGenerationStateFlow every { billingRepo.billingStateFlow() } returns billingStateFlow every { selectedModelProvider.getSelectedModelFlow() } returns selectedModelFlow coEvery { orpheusChatRepository.restoreOrStartNewSession() } returns true every { sessionStore.sessionFlow } returns sessionFlow every { statsigManager.checkGate(any()) } returns false every { userPrefsDataStoreManager.getThemeMode() } returns themeModeFlow every { processorFactory.create( any(), any(), any(), any(), ) } answers { MviProcessorFactory( analyticsManager = analyticsManager, statsigManager = statsigManager, loggerFactory = mockk(relaxed = true), ).create(firstArg(), secondArg(), thirdArg(), arg(3)) } subject = OrpheusChatVM( processorFactory = processorFactory, orpheusChatRepository = orpheusChatRepository, sessionStore = sessionStore, listeningSourceCache = listeningSourceCache, mediaManager = mediaManager, generationRepository = generationRepository, updateClipReactionUseCase = updateClipReactionUseCase, billingRepo = billingRepo, selectedModelProvider = selectedModelProvider, shareLinkRepository = shareLinkRepository, featureManager = featureManager, projectsControllerFactory = projectControllerFactory, orpheusAnalyticsManager = orpheusAnalyticsManager, getStringFromResourcesUseCase = getStringFromResourcesUseCase, userPrefsDataStoreManager = userPrefsDataStoreManager, ) } @After fun tearDown() { Dispatchers.resetMain() } @Test fun `given message text when MessageChanged then updates current message in state`() = runTest { val messageText = "Hello Orpheus" subject.sendEvent(OrpheusChatEvent.MessageChanged(message = messageText)) advanceUntilIdle() assertEquals(messageText, subject.state.value.currentMessage) } @Test fun `given empty message when SendMessage then state unchanged and no repository call`() = runTest { subject.sendEvent(OrpheusChatEvent.MessageChanged(message = " ")) subject.sendEvent(OrpheusChatEvent.SendMessage) advanceUntilIdle() coVerify(exactly = 0) { orpheusChatRepository.sendMessage(any()) } coVerify(exactly = 0) { orpheusChatRepository.createNewProjectInCurrentSession(any()) } assertEquals(" ", subject.state.value.currentMessage) } @Test fun `given non-empty message when SendMessage then sends message and clears input`() = runTest { val messageText = "Create a song about cats" subject.sendEvent(OrpheusChatEvent.MessageChanged(message = messageText)) subject.sendEvent(OrpheusChatEvent.SendMessage) advanceUntilIdle() coVerify { orpheusChatRepository.sendMessage(content = messageText) } assertEquals("", subject.state.value.currentMessage) } @Test fun `given message with whitespace when SendMessage then trims and sends message`() = runTest { val messageText = " Hello " subject.sendEvent(OrpheusChatEvent.MessageChanged(message = messageText)) subject.sendEvent(OrpheusChatEvent.SendMessage) advanceUntilIdle() coVerify { orpheusChatRepository.sendMessage(content = "Hello") } assertEquals("", subject.state.value.currentMessage) } @Test fun `given current message when NewChat then creates new session and clears message`() = runTest { subject.sendEvent(OrpheusChatEvent.MessageChanged(message = "Some text")) subject.sendEvent(OrpheusChatEvent.NewChat) advanceUntilIdle() coVerify { orpheusChatRepository.restoreOrStartNewSession() } coVerify { orpheusChatRepository.startNewSession() } assertEquals("", subject.state.value.currentMessage) } @Test fun `given clip not playing when ToggleClipPlayback then sets single playing clip`() = runTest { val clip = createClipData(clipId = "clip-1") subject.sendEvent(OrpheusChatEvent.ToggleClipPlayback(messageId = testMessageId, clip = clip)) advanceUntilIdle() coVerify { mediaManager.setSinglePlayingClipData(clip) } } @Test fun `given clip already playing when ToggleClipPlayback then pauses playback`() = runTest { val clipId = Id("clip-1") val clip = createClipData(clipId = clipId.value) subject.sendEvent(OrpheusChatEvent.ToggleClipPlayback(messageId = testMessageId, clip = clip)) advanceUntilIdle() mediaPlayerStateFlow.value = createMediaPlayerState( clipId = clipId, isPlaying = true, ) advanceUntilIdle() subject.sendEvent(OrpheusChatEvent.ToggleClipPlayback(messageId = testMessageId, clip = clip)) advanceUntilIdle() coVerify { mediaManager.setIsPlaying(false) } coVerify(exactly = 1) { mediaManager.setSinglePlayingClipData(clip) } } @Test fun `given clip playing but different clip when ToggleClipPlayback then plays new clip`() = runTest { val currentClipId = Id("clip-1") val newClip = createClipData(clipId = "clip-2") mediaPlayerStateFlow.value = createMediaPlayerState( clipId = currentClipId, isPlaying = true, ) advanceUntilIdle() subject.sendEvent(OrpheusChatEvent.ToggleClipPlayback(messageId = testMessageId, clip = newClip)) advanceUntilIdle() coVerify { mediaManager.setSinglePlayingClipData(newClip) } } @Test fun `given clip paused when ToggleClipPlayback then starts playing same clip`() = runTest { val clipId = Id("clip-1") val clip = createClipData(clipId = clipId.value) mediaPlayerStateFlow.value = createMediaPlayerState(clipId) advanceUntilIdle() subject.sendEvent(OrpheusChatEvent.ToggleClipPlayback(messageId = testMessageId, clip = clip)) advanceUntilIdle() coVerify { mediaManager.setSinglePlayingClipData(clip) } } @Test fun `given selected clip when ScrubClip with matching clip then calls mediaManager scrubToPercent`() = runTest { val clipId = Id("clip-1") val clip = createClipData(clipId = clipId.value) val scrubPercent = 0.5f subject.sendEvent(OrpheusChatEvent.SelectClip(messageId = testMessageId, clip = clip)) advanceUntilIdle() subject.sendEvent( OrpheusChatEvent.ScrubClip( messageId = testMessageId, clip = clip, scrubPercent = scrubPercent, ), ) advanceUntilIdle() verify { mediaManager.scrubToPercent(percent = scrubPercent) } } @Test fun `given no selected clip when ScrubClip then does not call mediaManager scrubToPercent`() = runTest { val clip = createClipData(clipId = "clip-1") subject.sendEvent( OrpheusChatEvent.ScrubClip( messageId = testMessageId, clip = clip, scrubPercent = 0.5f, ), ) advanceUntilIdle() verify(exactly = 0) { mediaManager.scrubToPercent(any()) } } @Test fun `given clip when ShareClip then shows share link bottom sheet`() = runTest { val clip = createClipData(clipId = "clip-1") subject.sendEvent(OrpheusChatEvent.ShareClip(messageId = testMessageId, clip = clip)) advanceUntilIdle() val bottomSheetState = subject.state.value.bottomSheetState assertTrue(bottomSheetState is OrpheusChatState.BottomSheetState.ShareLinkVisible) assertEquals(clip, (bottomSheetState as OrpheusChatState.BottomSheetState.ShareLinkVisible).clip) } @Test fun `given clip when OpenClipOverflow then shows song actions bottom sheet`() = runTest { val clip = createClipData(clipId = "clip-1") subject.sendEvent(OrpheusChatEvent.OpenClipOverflow(messageId = testMessageId, clip = clip)) advanceUntilIdle() val bottomSheetState = subject.state.value.bottomSheetState assertTrue(bottomSheetState is OrpheusChatState.BottomSheetState.SongActionsVisible) assertEquals(clip, (bottomSheetState as OrpheusChatState.BottomSheetState.SongActionsVisible).clip) } @Test fun `given song actions sheet shown when DismissBottomSheet then hides sheet`() = runTest { val clip = createClipData(clipId = "clip-1") subject.sendEvent(OrpheusChatEvent.OpenClipOverflow(messageId = testMessageId, clip = clip)) advanceUntilIdle() subject.sendEvent(OrpheusChatEvent.DismissBottomSheet) advanceUntilIdle() assertTrue(subject.state.value.bottomSheetState is OrpheusChatState.BottomSheetState.Hidden) } @Test fun `given share link bottom sheet shown when DismissBottomSheet then hides sheet`() = runTest { val clip = createClipData(clipId = "clip-1") subject.sendEvent(OrpheusChatEvent.ShareClip(messageId = testMessageId, clip = clip)) advanceUntilIdle() subject.sendEvent(OrpheusChatEvent.DismissBottomSheet) advanceUntilIdle() assertTrue(subject.state.value.bottomSheetState is OrpheusChatState.BottomSheetState.Hidden) } @Test fun `given messages updated when MessagesUpdated then updates state with messages`() = runTest { val message1 = createOrpheusMessage(messageId = Id("msg-1")) val message2 = createOrpheusMessage(messageId = Id("msg-2")) val messages = persistentListOf(message1, message2) val asyncData = AsyncData.Ready(messages) messagesFlow.value = asyncData advanceUntilIdle() assertEquals(asyncData, subject.state.value.messages) } @Test fun `given messages with generating clip when MessagesUpdated then adds clip ID to song generation state store`() = runTest { val clip1 = createClipData(clipId = "clip-1", status = ClipStatus.Submitted) val clip2 = createClipData(clipId = "clip-2") val message = createOrpheusMessage( messageId = Id("msg-1"), contentType = MessageContentType.GeneratedClips, clips = listOf(clip1, clip2), ) messagesFlow.value = AsyncData.Ready(persistentListOf(message)) advanceUntilIdle() verify { generationRepository.upsertClipStatuses( match { clipStatuses -> clipStatuses.size == 1 && clipStatuses[clip1.clipId] == ClipStatus.Submitted }, ) } } @Test fun `given messages without clips when MessagesUpdated then does not add to song generation manager`() = runTest { val message = createOrpheusMessage(messageId = Id("msg-1"), clips = emptyList()) messagesFlow.value = AsyncData.Ready(persistentListOf(message)) advanceUntilIdle() verify(exactly = 0) { generationRepository.upsertClipStatuses(any()) } } @Test fun `given media player state updated when MediaPlayerStateUpdated then updates playback state`() = runTest { val clipId = Id("clip-1") val clip = createClipData(clipId = clipId.value) val progress = 0.5f subject.sendEvent(OrpheusChatEvent.ToggleClipPlayback(messageId = testMessageId, clip = clip)) advanceUntilIdle() mediaPlayerStateFlow.value = createMediaPlayerState( clipId = clipId, isPlaying = true, progress = progress, ) advanceUntilIdle() val selectedClip = subject.state.value.selectedMessageClip assertNotNull(selectedClip) assertEquals(clipId, selectedClip!!.clip.clipId) assertTrue(selectedClip.isPlaying) assertEquals(progress, selectedClip.playbackProgress, 0.001f) } @Test fun `given no clip playing when MediaPlayerStateUpdated then sets null clip`() = runTest { mediaPlayerStateFlow.value = createMediaPlayerState() advanceUntilIdle() assertNull(subject.state.value.selectedMessageClip) } @Test fun `given ready clip IDs when GeneratedClipsReady then updates clips status in repository`() = runTest { val clipId1 = Id("clip-1") val clipId2 = Id("clip-2") songGenerationStateFlow.value = SongGenerationStoreImpl.SongGenerationState( clipStatuses = mapOf( clipId1 to ClipStatus.Streaming, clipId2 to ClipStatus.Complete, ).toImmutableMap(), ) advanceUntilIdle() coVerify { orpheusChatRepository.updateClipsStatusInMessages( match { ids -> ids.contains(clipId1) && ids.contains(clipId2) && ids.size == 2 }, ) } } @Test fun `given valid billing info when model selection updates then registers correct external key`() = runTest { val model1 = createModel(name = "v3.5", externalKey = "chirp-v3-5-key") val model2 = createModel(name = "v5", externalKey = "chirp-crow-key") val billingInfo = createBillingInfo(models = listOf(model1, model2)) billingStateFlow.value = billingInfo selectedModelFlow.value = model2 advanceUntilIdle() coVerify { orpheusChatRepository.registerSelectedModel(modelName = "chirp-crow-key") } } @Test fun `given null billing info when model selection updates then registers fallback model`() = runTest { billingStateFlow.value = null selectedModelFlow.value = SelectedModelProvider.FALLBACK_DEFAULT_MODEL advanceUntilIdle() coVerify { orpheusChatRepository.registerSelectedModel(modelName = "chirp-auk-turbo") } } @Test fun `given empty models list when model selection updates then registers fallback model`() = runTest { val billingInfo = createBillingInfo(models = emptyList()) billingStateFlow.value = billingInfo selectedModelFlow.value = SelectedModelProvider.FALLBACK_DEFAULT_MODEL advanceUntilIdle() coVerify { orpheusChatRepository.registerSelectedModel(modelName = "chirp-auk-turbo") } } @Test fun `given model selection changes when model selection updates then registers new model`() = runTest { val model1 = createModel(name = "v3.5", externalKey = "chirp-v3-5-key") val model2 = createModel(name = "v5", externalKey = "chirp-crow-key") val billingInfo = createBillingInfo(models = listOf(model1, model2)) billingStateFlow.value = billingInfo selectedModelFlow.value = model1 advanceUntilIdle() coVerify { orpheusChatRepository.registerSelectedModel(modelName = "chirp-v3-5-key") } selectedModelFlow.value = model2 advanceUntilIdle() coVerify { orpheusChatRepository.registerSelectedModel(modelName = "chirp-crow-key") } } @Test fun `given billing info changes when model selection updates then registers based on new info`() = runTest { val model1v1 = createModel(name = "v5", externalKey = "chirp-crow-key-v1") val billingInfoV1 = createBillingInfo(models = listOf(model1v1)) billingStateFlow.value = billingInfoV1 selectedModelFlow.value = model1v1 advanceUntilIdle() coVerify { orpheusChatRepository.registerSelectedModel(modelName = "chirp-crow-key-v1") } val model1v2 = createModel(name = "v5", externalKey = "chirp-crow-key-v2") val billingInfoV2 = createBillingInfo(models = listOf(model1v2)) billingStateFlow.value = billingInfoV2 selectedModelFlow.value = model1v2 advanceUntilIdle() coVerify { orpheusChatRepository.registerSelectedModel(modelName = "chirp-crow-key-v2") } } @Test fun `given valid billing info when model selection updates then updates state with external key`() = runTest { val model1 = createModel(name = "v3.5", externalKey = "chirp-v3-5-key") val model2 = createModel(name = "v5", externalKey = "chirp-crow-key") val billingInfo = createBillingInfo(models = listOf(model1, model2)) billingStateFlow.value = billingInfo selectedModelFlow.value = model2 advanceUntilIdle() assertEquals("chirp-crow-key", subject.state.value.selectedModel?.externalKey) } @Test fun `given null billing info when model selection updates then updates state with fallback model`() = runTest { billingStateFlow.value = null selectedModelFlow.value = SelectedModelProvider.FALLBACK_DEFAULT_MODEL advanceUntilIdle() assertEquals("chirp-auk-turbo", subject.state.value.selectedModel?.externalKey) } @Test fun `given model changes when model selection updates then updates state with new external key`() = runTest { val model1 = createModel(name = "v3.5", externalKey = "chirp-v3-5-key") val model2 = createModel(name = "v5", externalKey = "chirp-crow-key") val billingInfo = createBillingInfo(models = listOf(model1, model2)) billingStateFlow.value = billingInfo selectedModelFlow.value = model1 advanceUntilIdle() assertEquals("chirp-v3-5-key", subject.state.value.selectedModel?.externalKey) selectedModelFlow.value = model2 advanceUntilIdle() assertEquals("chirp-crow-key", subject.state.value.selectedModel?.externalKey) } @Test fun `given credit count changes when billing info updates then updates state with new credit count`() = runTest { val billingInfo1 = createBillingInfo(totalCreditsLeft = 1000) val billingInfo2 = createBillingInfo(totalCreditsLeft = 500) billingStateFlow.value = billingInfo1 advanceUntilIdle() assertEquals(1000, subject.state.value.creditCount) billingStateFlow.value = billingInfo2 advanceUntilIdle() assertEquals(500, subject.state.value.creditCount) } @Test fun `given CurrentProjectChanged event when project provided then updates state with current project`() = runTest { val project = object : BaseProjectMetadata { override val id = "project-123".asId() override val name = "New Project" } subject.sendEvent(OrpheusChatEvent.Internal.CurrentProjectChanged(project = project)) advanceUntilIdle() assertEquals(project, subject.state.value.currentProject) } @Test fun `given session with project when session flow updates then state contains current project`() = runTest { val project = object : BaseProjectMetadata { override val id = "new-project".asId() override val name = "Updated Project" } val newSession = testSession.copy(project = project) sessionFlow.value = AsyncData.Ready(newSession) advanceUntilIdle() assertEquals(project, subject.state.value.currentProject) } @Test fun `given CreateProject event when triggered then calls startNewSession and emits CloseProjectsDrawer effect`() = runTest { coEvery { orpheusChatRepository.startNewSession() } returns Either.Right( mockk(relaxed = true), ) subject.effects.test { subject.sendEvent(OrpheusChatEvent.Internal.CreateProject) advanceUntilIdle() val effect = awaitItem() assertTrue(effect is OrpheusChatEffect.CloseProjectsDrawer) coVerify { orpheusChatRepository.startNewSession() } cancelAndIgnoreRemainingEvents() } } @Test fun `given valid message when SendMessage then sets isWaitingForResponse true and clears message`() = runTest { val messageText = "Create a song" subject.sendEvent(OrpheusChatEvent.MessageChanged(message = messageText)) subject.sendEvent(OrpheusChatEvent.SendMessage) advanceUntilIdle() assertEquals(true, subject.state.value.isWaitingForResponse) assertEquals("", subject.state.value.currentMessage) } @Test fun `given waiting for response when new assistant message then sets isWaitingForResponse false`() = runTest { subject.sendEvent(OrpheusChatEvent.MessageChanged(message = "Hello")) subject.sendEvent(OrpheusChatEvent.SendMessage) advanceUntilIdle() assertEquals(true, subject.state.value.isWaitingForResponse) val assistantMessage = createOrpheusMessage(messageId = Id("msg-1")) messagesFlow.value = AsyncData.Ready(persistentListOf(assistantMessage)) advanceUntilIdle() assertEquals(false, subject.state.value.isWaitingForResponse) } @Test fun `given waiting for response when no assistant message then keeps isWaitingForResponse true`() = runTest { subject.sendEvent(OrpheusChatEvent.MessageChanged(message = "Hello")) subject.sendEvent(OrpheusChatEvent.SendMessage) advanceUntilIdle() assertEquals(true, subject.state.value.isWaitingForResponse) val userMessage = createOrpheusMessage(messageId = Id("msg-1")) messagesFlow.value = AsyncData.Ready(persistentListOf(userMessage.copy(role = MessageRole.User))) advanceUntilIdle() assertEquals(true, subject.state.value.isWaitingForResponse) } @Test fun `given isWaitingForResponse true when StreamingError then sets isWaitingForResponse to false`() = runTest { subject.sendEvent(OrpheusChatEvent.MessageChanged(message = "Hello")) subject.sendEvent(OrpheusChatEvent.SendMessage) advanceUntilIdle() assertEquals(true, subject.state.value.isWaitingForResponse) subject.sendEvent(OrpheusChatEvent.Internal.StreamingError(error = "Connection failed")) advanceUntilIdle() assertEquals(false, subject.state.value.isWaitingForResponse) } @Test fun `given not waiting when new assistant message then keeps isWaitingForResponse false`() = runTest { assertEquals(false, subject.state.value.isWaitingForResponse) val assistantMessage = createOrpheusMessage(messageId = Id("msg-1")) messagesFlow.value = AsyncData.Ready(persistentListOf(assistantMessage)) advanceUntilIdle() assertEquals(false, subject.state.value.isWaitingForResponse) } @Test fun `given waiting with multiple messages when last is assistant then sets false`() = runTest { subject.sendEvent(OrpheusChatEvent.MessageChanged(message = "Hello")) subject.sendEvent(OrpheusChatEvent.SendMessage) advanceUntilIdle() assertEquals(true, subject.state.value.isWaitingForResponse) val userMessage = createOrpheusMessage(messageId = Id("msg-1"), role = MessageRole.User) val assistantMessage1 = createOrpheusMessage(messageId = Id("msg-2"), role = MessageRole.Assistant) val assistantMessage2 = createOrpheusMessage(messageId = Id("msg-3"), role = MessageRole.Assistant) messagesFlow.value = AsyncData.Ready(persistentListOf(userMessage, assistantMessage1, assistantMessage2)) advanceUntilIdle() assertEquals(false, subject.state.value.isWaitingForResponse) } @Test fun `given waiting with multiple messages when last not assistant then keeps true`() = runTest { subject.sendEvent(OrpheusChatEvent.MessageChanged(message = "Hello")) subject.sendEvent(OrpheusChatEvent.SendMessage) advanceUntilIdle() assertEquals(true, subject.state.value.isWaitingForResponse) val assistantMessage = createOrpheusMessage(messageId = Id("msg-1"), role = MessageRole.Assistant) val userMessage = createOrpheusMessage(messageId = Id("msg-2"), role = MessageRole.User) messagesFlow.value = AsyncData.Ready(persistentListOf(assistantMessage, userMessage)) advanceUntilIdle() assertEquals(true, subject.state.value.isWaitingForResponse) } @Test fun `given Link platform and successful API call when StartShare then emits effect and dismisses bottom sheet`() = runTest { val clip = createClipData(clipId = "clip-1") val platform = SharePlatformConstants.Link.Copy val shareUrl = Url("https://suno.com/song/clip-1") coEvery { shareLinkRepository.getSongShareLink( contentId = clip.clipId, platform = platform.backendValue, ) } returns Either.Right(shareUrl) subject.effects.test { subject.sendEvent( OrpheusChatEvent.StartShare( messageId = testMessageId, clip = clip, platform = platform, ), ) advanceUntilIdle() val effect = awaitItem() as OrpheusChatEffect.ShareLink assertEquals(clip.clipId, effect.clipId) assertEquals(platform, effect.sharePlatform) assertEquals(shareUrl, effect.link) advanceUntilIdle() assertTrue(subject.state.value.bottomSheetState is OrpheusChatState.BottomSheetState.Hidden) } } @Test fun `given Link platform and API error when StartShare then does not emit effect and keeps bottom sheet open`() = runTest { val clip = createClipData(clipId = "clip-1") val platform = SharePlatformConstants.Link.WhatsApp val error = mockk(relaxed = true) coEvery { shareLinkRepository.getSongShareLink( contentId = clip.clipId, platform = platform.backendValue, ) } returns Either.Left(error) subject.sendEvent(OrpheusChatEvent.ShareClip(messageId = testMessageId, clip = clip)) advanceUntilIdle() subject.effects.test { subject.sendEvent( OrpheusChatEvent.StartShare( messageId = testMessageId, clip = clip, platform = platform, ), ) advanceUntilIdle() expectNoEvents() assertTrue(subject.state.value.bottomSheetState is OrpheusChatState.BottomSheetState.ShareLinkVisible) } } @Test fun `given song when SongDeleted then removes clip from media manager and repository`() = runTest { val song = createSongListData(id = "clip-1") subject.sendEvent( OrpheusChatEvent.SongDeleted( messageId = testMessageId, song = song, ), ) advanceUntilIdle() verify { mediaManager.removeClipById(song.id) } verify { orpheusChatRepository.removeClipFromMessage( messageId = testMessageId, clipId = song.id, ) } } @Test fun `given song when SongRenamed then updates clip in media manager and repository with new title`() = runTest { val song = createSongListData(id = "clip-1", title = "Updated Song Title") subject.sendEvent( OrpheusChatEvent.SongRenamed( messageId = testMessageId, song = song, ), ) advanceUntilIdle() verify { mediaManager.updateClip( oldLocalClipData = match { it.clipId == song.id }, newLocalClipData = match { it.nowPlayingTitle == song.title }, ) } verify { orpheusChatRepository.updateClipInMessage( messageId = testMessageId, updatedClip = match { it.clipId == song.id && it.nowPlayingTitle == song.title }, ) } } private fun createClipData( clipId: String, status: ClipStatus = ClipStatus.Complete, reaction: ReactionType = ReactionType.NOTHING, ): LocalClipData = LocalClipData( clipId = Id(clipId), mediaType = SunoMediaType.AUDIO, mediaUrl = Url("https://example.com/audio.mp3"), artistName = "Test Artist", artistUserId = Id("user-123"), handle = UserHandle("testuser"), artistAvatarUrl = Url("https://example.com/avatar.jpg"), nowPlayingTitle = "Test Song", albumImageUrl = Url("https://example.com/album.jpg"), videoCoverUrl = null, caption = "Test caption", isPublic = true, status = status, upvoteCount = 10, reaction = reaction, modelName = "chirp-v3", majorModelVersion = "v3", prompt = null, gptPrompt = null, tags = null, displayTags = null, commentCount = 0, downloadDisabledReason = null, previewUrl = Url("https://example.com/preview.mp3"), playCount = 100, isFollowing = false, ) private fun createOrpheusMessage( messageId: Id, clips: List = emptyList(), role: MessageRole = MessageRole.Assistant, contentType: MessageContentType = MessageContentType.Chat, ): OrpheusMessage = OrpheusMessage( messageId = messageId, role = role, content = "Test message content", timestamp = Duration.ZERO, generatedClips = clips, contentType = contentType, status = AssistantMessageStatus.Streaming, ) private fun createMediaPlayerState( clipId: Id? = null, isPlaying: Boolean = false, progress: Float = 0f, ): MediaPlayerState = MediaPlayerState( localClipDataQueue = if (clipId != null) { listOf(createClipData(clipId = clipId.value)) } else { emptyList() }, nowPlayingClipIndex = if (clipId != null) 0 else null, isPlaying = isPlaying, percentageComplete = progress, repeatMode = RepeatMode.REPEAT_MODE_NONE, playtimeDuration = Duration.ZERO, ) 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(), totalCreditsLeft: Int? = null, ): SubscriptionInfoResponse = mockk(relaxed = true) { every { this@mockk.models } returns models every { this@mockk.totalCreditsLeft } returns (totalCreditsLeft ?: 0) } @Test fun `given message with clips when CreateMoreClips then calls generation API with correct parameters`() = runTest { val model = createModel(name = "chirp-crow", externalKey = "chirp-crow-key") billingStateFlow.value = createBillingInfo(models = listOf(model)) advanceUntilIdle() val clip = createClipData( clipId = "clip-1", status = ClipStatus.Complete, ) val message = createOrpheusMessage( messageId = testMessageId, clips = listOf(clip), contentType = MessageContentType.GeneratedClips, ) val generatedClipIds = setOf(Id("new-clip-1"), Id("new-clip-2")) coEvery { generationRepository.startSongGeneration(any()) } returns Either.Right( mockk(relaxed = true) { every { clips } returns generatedClipIds.map { mockk { every { id } returns it.value } } }, ) subject.sendEvent(OrpheusChatEvent.CreateMoreClips(message = message)) advanceUntilIdle() coVerify { generationRepository.startSongGeneration( match { params -> params.generationType.name == "TEXT" && !params.makeInstrumental }, ) } assertNull(subject.state.value.createMoreClipsMessageId) } @Test fun `given message with no clips when CreateMoreClips then does not set createMoreClipsMessageId or call API`() = runTest { val message = createOrpheusMessage( messageId = testMessageId, clips = emptyList(), ) subject.sendEvent(OrpheusChatEvent.CreateMoreClips(message = message)) advanceUntilIdle() assertNull(subject.state.value.createMoreClipsMessageId) coVerify(exactly = 0) { generationRepository.startSongGeneration(any()) } } @Test fun `given create more clips success when clips created then adds clips and clears state`() = runTest { val model = createModel(name = "chirp-crow", externalKey = "chirp-crow-key") billingStateFlow.value = createBillingInfo(models = listOf(model)) advanceUntilIdle() val clip = createClipData(clipId = "clip-1") val message = createOrpheusMessage( messageId = testMessageId, clips = listOf(clip), contentType = MessageContentType.GeneratedClips, ) val generatedClipIds = setOf(Id("new-clip-1"), Id("new-clip-2")) coEvery { generationRepository.startSongGeneration(any()) } returns Either.Right( mockk(relaxed = true) { every { clips } returns generatedClipIds.map { mockk { every { id } returns it.value } } }, ) subject.sendEvent(OrpheusChatEvent.CreateMoreClips(message = message)) advanceUntilIdle() coVerify { orpheusChatRepository.fetchAndUpdateMessageClips( messageId = testMessageId, clipIds = generatedClipIds, replaceClips = false, ) } assertNull(subject.state.value.createMoreClipsMessageId) } @Test fun `given create more clips fails when API error then clears createMoreClipsMessageId without adding clips`() = runTest { val clip = createClipData(clipId = "clip-1") val message = createOrpheusMessage( messageId = testMessageId, clips = listOf(clip), contentType = MessageContentType.GeneratedClips, ) coEvery { generationRepository.startSongGeneration(any()) } returns Either.Left(mockk(relaxed = true)) subject.sendEvent(OrpheusChatEvent.CreateMoreClips(message = message)) advanceUntilIdle() coVerify(exactly = 0) { orpheusChatRepository.fetchAndUpdateMessageClips(any(), any(), any()) } assertNull(subject.state.value.createMoreClipsMessageId) } @Test fun `given message with clip data when CreateMoreClips then extracts correct generation parameters`() = runTest { val model = createModel(name = "v5", externalKey = "chirp-crow-key") val billingInfo = createBillingInfo(models = listOf(model)) billingStateFlow.value = billingInfo selectedModelFlow.value = model advanceUntilIdle() val clip = createClipData( clipId = "clip-1", status = ClipStatus.Complete, ).copy( prompt = "Test lyrics", tags = "rock, energetic", nowPlayingTitle = "Test Song", modelName = "chirp-v4", gptPrompt = "Test gpt prompt", ) val message = createOrpheusMessage( messageId = testMessageId, clips = listOf(clip), contentType = MessageContentType.GeneratedClips, ) coEvery { generationRepository.startSongGeneration(any()) } returns Either.Right( mockk(relaxed = true) { every { clips } returns listOf(mockk { every { id } returns "new-clip-1" }) }, ) subject.sendEvent(OrpheusChatEvent.CreateMoreClips(message = message)) advanceUntilIdle() coVerify { generationRepository.startSongGeneration( match { params -> params.prompt == "Test lyrics" && params.tags == "rock, energetic" && params.title == "Test Song" && params.modelVersionName == "chirp-crow-key" && params.gptPrompt == "Test gpt prompt" }, ) } } @Test fun `given NewChat event and session restoration succeeds when startNewSession then registers selected model`() = runTest { val model = createModel(name = "v5", externalKey = "chirp-crow-key") val billingInfo = createBillingInfo(models = listOf(model)) billingStateFlow.value = billingInfo selectedModelFlow.value = model coEvery { orpheusChatRepository.startNewSession().isRight() } returns true advanceUntilIdle() subject.sendEvent(OrpheusChatEvent.NewChat) advanceUntilIdle() coVerify { orpheusChatRepository.startNewSession() } coVerify { orpheusChatRepository.registerSelectedModel(modelName = "chirp-crow-key") } } private fun createSongListData( id: String, title: String = "Test Song", ): SongListData = SongListData( id = Id(id), createdAt = Date(), status = ClipStatus.Complete, displayImageUrl = Url("https://example.com/image.jpg"), mediaUrl = Url("https://example.com/audio.mp3"), videoCoverUrl = null, previewUrl = Url("https://example.com/preview.mp3"), artistName = "Test Artist", artistUserId = Id("user-123"), handle = UserHandle("testuser"), artistAvatarUrl = Url("https://example.com/avatar.jpg"), title = title, tags = null, displayTags = null, playCount = 100, isPublic = true, modelName = "chirp-v3", majorModelVersion = "v3", reaction = null, prompt = null, gptPrompt = null, commentCount = 0, upvoteCount = 10, caption = null, captionMentions = null, downloadDisabledReason = null, duration = 120.seconds, canRemix = true, isRemix = false, remixTask = null, optOutVideoCoverHook = false, type = null, ) @Test fun `given valid project when ShareProject and fetch succeeds then calls repo and emits ShareChatLink effect`() = runTest { val projectId = Id("project-1") val project = createProjectMetadata(id = projectId) val shareUrl = Url("https://suno.com/chat/project-1") coEvery { orpheusChatRepository.getChatLinkForProject(projectId) } returns Either.Right(shareUrl) subject.effects.test { subject.sendEvent(OrpheusChatEvent.ShareProject(project = project)) advanceUntilIdle() val effect = awaitItem() as OrpheusChatEffect.ShareChatLink assertEquals(shareUrl, effect.link) coVerify { orpheusChatRepository.getChatLinkForProject(projectId) } cancelAndIgnoreRemainingEvents() } } @Test fun `given valid project when ShareProject and fetch fails then emits ShowSnackbar effect`() = runTest { val projectId = Id("project-1") val project = createProjectMetadata(id = projectId) coEvery { orpheusChatRepository.getChatLinkForProject(projectId) } returns Either.Left(Unit) subject.effects.test { subject.sendEvent(OrpheusChatEvent.ShareProject(project = project)) advanceUntilIdle() val effect = awaitItem() as OrpheusChatEffect.ShowSnackbar assertTrue(effect.message is UiString.Resource) cancelAndIgnoreRemainingEvents() } } @Test fun `given theme mode flow emits DARK when ThemeModeChanged then updates state with DARK theme`() = runTest { themeModeFlow.value = ThemeMode.DARK advanceUntilIdle() assertEquals(ThemeMode.DARK, subject.state.value.themeMode) } private fun createProjectMetadata( id: Id, ): ProjectMetadata = mockk { every { this@mockk.id } returns id every { name } returns "Test Project" every { description } returns "Test Description" every { clipCount } returns 5 every { lastUpdatedClipTime } returns null every { createdAt } returns null } }