package com.suno.android.common_data.repos import arrow.core.Either import arrow.core.getOrElse import arrow.retrofit.adapter.either.networkhandling.HttpError import com.suno.android.common_core_utils.ApplicationCoroutineScope import com.suno.android.common_core_utils.Id import com.suno.android.common_core_utils.SunoLogger import com.suno.android.common_core_utils.asId import com.suno.android.common_data.generation.AlignedLyrics import com.suno.android.common_data.generation.LyricChunk import com.suno.android.common_data.generation.SongGenerationStateStore 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.SongListData import com.suno.android.common_data.mappers.share_asset.ShareAssetStatus import com.suno.android.common_data.mappers.share_asset.asShareAssetStatusOrNull import com.suno.android.common_data.user.UserSessionRepository import com.suno.android.common_networking.extensions.ErrorResponse import com.suno.android.common_networking.extensions.toThrowable import com.suno.android.common_networking.generation.GenerationRealtimeClient import com.suno.android.common_networking.remote.common.RemoteClip import com.suno.android.common_networking.remote.entities.GenParamsSpec import com.suno.android.common_networking.remote.entities.GenerationRequestSchema import com.suno.android.common_networking.remote.entities.ImagePromptSpec import com.suno.android.common_networking.remote.entities.RemoteShareAssetBody import com.suno.android.common_networking.remote.feed.FeedService import com.suno.android.common_networking.remote.gen.GenService import com.suno.android.common_networking.remote.generate.GenerateService import com.suno.android.common_networking.remote.session.User import com.suno.android.common_networking.sse.RealtimeEvent import com.suno.android.gating.Feature import com.suno.android.gating.FeatureManager import kotlinx.coroutines.CoroutineScope import kotlinx.coroutines.ExperimentalCoroutinesApi import kotlinx.coroutines.delay import kotlinx.coroutines.flow.Flow import kotlinx.coroutines.flow.SharingStarted import kotlinx.coroutines.flow.StateFlow import kotlinx.coroutines.flow.WhileSubscribed import kotlinx.coroutines.flow.distinctUntilChanged import kotlinx.coroutines.flow.emptyFlow import kotlinx.coroutines.flow.flatMapLatest import kotlinx.coroutines.flow.flow import kotlinx.coroutines.flow.flowOf import kotlinx.coroutines.flow.last import kotlinx.coroutines.flow.map import kotlinx.coroutines.flow.onEach import kotlinx.coroutines.flow.shareIn import kotlinx.serialization.json.Json import javax.inject.Inject import kotlin.math.roundToInt import kotlin.time.Duration import kotlin.time.Duration.Companion.seconds import kotlin.time.DurationUnit typealias ShareAsset = Unit interface GenerationRepository { suspend fun getAlignedLyrics( clipId: Id, ): Result suspend fun startShareAssetGeneration( clipId: Id, startTime: Duration, endTime: Duration, ): Result suspend fun getShareAssetStatus( clipId: Id, shareAssetId: Id, ): Result suspend fun promptSongImage( prompt: String, ): Result suspend fun startSongGeneration( spec: GenParamsSpec, ): Either fun pollAllGeneratingSongs(): Flow>> fun songGenerationStateFlow(): StateFlow fun upsertClipStatuses( clipStatuses: Map, ClipStatus>, ) } @OptIn(ExperimentalCoroutinesApi::class) internal class DefaultGenerationRepository @Inject constructor( loggerFactory: SunoLogger.Factory, private val genService: GenService, private val generateService: GenerateService, @ApplicationCoroutineScope private val applicationScope: CoroutineScope, private val songGenerationStateStore: SongGenerationStateStore, private val feedService: FeedService, private val jsonParser: Json, private val generationRealtimeClient: GenerationRealtimeClient, private val userSessionRepository: UserSessionRepository, private val featureManager: FeatureManager, ) : GenerationRepository { private val logger = loggerFactory.create(this@DefaultGenerationRepository) private val pollingIntervalSeconds by lazy { featureManager.getFeatureValue(Feature.SongGenerationPollingIntervalSeconds).value.roundToInt() } private val generationUpdatesFlow = songGenerationStateStore.songGenerationStateFlow() .map { it.nonTerminalClipIds } .distinctUntilChanged() .flatMapLatest { songIds -> if (songIds.isEmpty()) return@flatMapLatest emptyFlow() if (featureManager.hasFeature(Feature.RealtimeSongGeneration)) { userSessionRepository.sessionConfigurationStateFlow() .flatMapLatest { sessionConfig -> val userId = sessionConfig.user?.id?.asId() if (userId != null) { logger.d { "Using real-time updates for user ${userId.value} with ${songIds.size} songs" } createRealtimeFlow(userId = userId, songIds = songIds) } else { logger.d { "No user ID, falling back to polling for ${songIds.size} songs" } createPollingFlow(songIds) } } } else { logger.d { "Real-time disabled, using HTTP polling for ${songIds.size} songs" } createPollingFlow(songIds) } } .onEach { either -> either.onRight { clips -> logger.d { "Received generation update with ${clips.size} clips" } val clipStatuses = clips.associate { remoteClip -> Id(remoteClip.id) to ClipStatus.fromString(remoteClip.status) } songGenerationStateStore.upsertClips(clipStatuses) } } .shareIn( scope = applicationScope, started = SharingStarted.WhileSubscribed(pollingIntervalSeconds.seconds / 2), replay = 1, ) private fun createRealtimeFlow( userId: Id, songIds: Set>, ): Flow>> = generationRealtimeClient.realtimeEventsFlow(userId) .flatMapLatest { event -> when (event) { is RealtimeEvent.Connected -> { logger.d { "Real-time generation updates connected" } emptyFlow() } is RealtimeEvent.Disconnected -> { logger.w { "Real-time disconnected, falling back to polling" } createPollingFlow(songIds) } is RealtimeEvent.Error -> { logger.e(event.throwable) { "Real-time error, falling back to polling" } createPollingFlow(songIds) } is RealtimeEvent.Data -> { flowOf(Either.Right(event.data)) } } } private fun createPollingFlow( songIds: Set>, ): Flow>> = flow { logger.d { "Starting polling with interval: ${pollingIntervalSeconds.seconds}" } while (true) { val response = feedService.getSongsWithIds(songIds.joinToString(",") { it.value }) response.onRight { feed -> emit(Either.Right(feed.clips)) }.onLeft { error -> logger.w { "Polling failed, will retry in ${pollingIntervalSeconds.seconds}: $error" } } delay(pollingIntervalSeconds.seconds) } } override suspend fun getAlignedLyrics( clipId: Id, ): Result { val result = genService.getAlignedLyrics(clipId.value).getOrElse { error -> val throwable = error.toThrowable() logger.e(throwable) return Result.failure(throwable) } return runCatching { AlignedLyrics( waveFormData = result.waveFormData!!, isStreamed = result.isStreamed!!, alignedWords = result.alignedWords?.map { remoteLyricChunk -> LyricChunk( word = remoteLyricChunk.word, startTime = remoteLyricChunk.startSecond.seconds, endTime = remoteLyricChunk.endSecond.seconds, pAlign = remoteLyricChunk.pAlign.toFloat(), ) } ?: emptyList(), ) } } override suspend fun startShareAssetGeneration( clipId: Id, startTime: Duration, endTime: Duration, ): Result = runCatching { genService.startShareAssetGeneration( genId = clipId.value, body = RemoteShareAssetBody( config = RemoteShareAssetBody.Config( presetId = "cover", presetStyle = "video", lyricsStyle = "lyrics_box", stickerStyle = "core_lyrics_standard", ), clipStartTimeSeconds = startTime.toDouble(DurationUnit.SECONDS), clipEndTimeSeconds = endTime.toDouble(DurationUnit.SECONDS), ), ).let { it.asShareAssetStatusOrNull() ?: error("failed to parse ShareAssetStatus") } } override suspend fun getShareAssetStatus( clipId: Id, shareAssetId: Id, ): Result = runCatching { genService.getShareAssetGenerationStatus( genId = clipId.value, assetId = shareAssetId.value, ).let { it.asShareAssetStatusOrNull(shareAssetId)!! } } override suspend fun promptSongImage( prompt: String, ): Result = runCatching { genService.promptImage(ImagePromptSpec(prompt)).last().body()?.imageUrl.toString() } override suspend fun startSongGeneration( spec: GenParamsSpec, ): Either = generateService.generateApiRunGenerationV2(spec) .mapLeft { val errorBody = (it as? HttpError)?.body runCatching { errorBody?.let { jsonParser.decodeFromString(it) } }.getOrNull() }.onRight { response -> val clipStatuses = response.clips.associate { clip -> Id(clip.id) to ClipStatus.fromString(clip.status) } songGenerationStateStore.upsertClips(clipStatuses) } override fun songGenerationStateFlow(): StateFlow = songGenerationStateStore.songGenerationStateFlow() override fun upsertClipStatuses( clipStatuses: Map, ClipStatus>, ) { songGenerationStateStore.upsertClips(clipStatuses) } override fun pollAllGeneratingSongs(): Flow>> = generationUpdatesFlow.map { it.mapLeft { } } }