package com.suno.android.common_data.upload import android.content.Context import androidx.core.net.toUri import androidx.hilt.work.HiltWorker import androidx.work.CoroutineWorker import androidx.work.WorkerParameters import androidx.work.workDataOf import arrow.core.getOrElse import com.suno.android.common_core_utils.Id import com.suno.android.common_core_utils.SunoLogger import com.suno.android.common_core_utils.model.Url import com.suno.android.common_data.mappers.clips.LocalClipData import com.suno.android.common_data.upload.model.UploadingFile import com.suno.android.common_networking.extensions.toThrowable import com.suno.android.common_networking.remote.aws.AwsUploadService import com.suno.android.common_networking.remote.entities.ClipMetadataSpec import com.suno.android.common_networking.remote.entities.CountingRequestBody import com.suno.android.common_networking.remote.entities.FinishUploadSpec import com.suno.android.common_networking.remote.entities.InputStreamRequestBody import com.suno.android.common_networking.remote.entities.UploadClipSchema import com.suno.android.common_networking.remote.entities.UploadRequestSpec import com.suno.android.common_networking.remote.entities.UploadRequestStatusSchema import com.suno.android.common_networking.remote.gen.GenService import com.suno.android.common_networking.remote.uploads.UploadsService import dagger.assisted.Assisted import dagger.assisted.AssistedInject import kotlinx.coroutines.Job import kotlinx.coroutines.coroutineScope import kotlinx.coroutines.delay import kotlinx.coroutines.launch import kotlinx.coroutines.plus import okhttp3.MediaType.Companion.toMediaType import okhttp3.MultipartBody import kotlin.math.pow import kotlin.time.Duration import kotlin.time.Duration.Companion.seconds import kotlin.time.times @HiltWorker class UploadFileAsClipWorker @AssistedInject constructor( loggerFactory: SunoLogger.Factory, @Assisted context: Context, @Assisted params: WorkerParameters, private val uploadsService: UploadsService, private val awsUploadService: AwsUploadService, private val genService: GenService, ) : CoroutineWorker(context, params) { object Parameters { object Input { const val FILE_URI = "file_uri" const val UPLOAD_NAME = "upload_name" } object Output { const val UPLOAD_ID = "upload_id" const val CLIP_ID = "clip_id" const val ERROR = "error" const val NAME = "name" const val THUMBNAIL = "thumbnail" } object Progress { const val UPLOAD_ID = "upload_id" const val STEP = "step" const val PERCENT = "percent" const val NAME = "name" const val THUMBNAIL = "thumbnail" enum class Steps { Initializing, Uploading, Finalizing, } } } private val logger = loggerFactory.create(this@UploadFileAsClipWorker) private suspend fun setProgress( uploadId: Id? = null, name: String, thumbnail: Url? = null, step: Parameters.Progress.Steps, progress: Float? = null, ) = setProgress( workDataOf( Parameters.Progress.STEP to step.name, Parameters.Progress.UPLOAD_ID to uploadId?.value, Parameters.Progress.NAME to name, Parameters.Progress.THUMBNAIL to thumbnail?.url, Parameters.Progress.PERCENT to progress, ), ) override suspend fun doWork(): Result = coroutineScope { val fileUri = inputData.getString(Parameters.Input.FILE_URI) ?: return@coroutineScope Result.failure( workDataOf( Parameters.Output.ERROR to "no ${Parameters.Input.FILE_URI} provided in inputData", ), ) val uploadName = inputData.getString(Parameters.Input.UPLOAD_NAME) ?: return@coroutineScope Result.failure( workDataOf( Parameters.Output.ERROR to "no ${Parameters.Input.UPLOAD_NAME} provided in inputData", ), ) val uri = fileUri.toUri() setProgress( name = uploadName, step = Parameters.Progress.Steps.Initializing, ) val uploadDestination = withRetry( retryCount = 5, minBackoffDelay = 2.seconds, ) { uploadsService.requestInfoForAudioUpload( UploadRequestSpec(extension = "wav"), ).getOrElse { error -> logger.e(error.toThrowable()) throw error.toThrowable() // caught in withRetry } }.getOrElse { error -> return@coroutineScope Result.failure( workDataOf( Parameters.Output.ERROR to "Failed to create upload", ), ) } val uploadId = Id(uploadDestination.id) var progressJob = Job() setProgress( uploadId = uploadId, name = uploadName, step = Parameters.Progress.Steps.Uploading, ) withRetry( retryCount = 5, minBackoffDelay = 2.seconds, ) { val inputStream = applicationContext.contentResolver.openInputStream(uri) ?: error("failed to open upload file InputStream") inputStream.use { awsUploadService.uploadFile( fields = uploadDestination.fields, file = MultipartBody.Part.createFormData( name = "file", filename = uri.path, body = CountingRequestBody( InputStreamRequestBody( contentType = "media/wav".toMediaType(), inputStream = inputStream, ), ) { bytesWritten, contentLength -> // drop progress updates if update scope is busy if (progressJob.isActive) return@CountingRequestBody val newJob = Job() progressJob = newJob val updateScope = (this@coroutineScope + newJob) updateScope.launch { setProgress( uploadId = uploadId, name = uploadName, step = Parameters.Progress.Steps.Uploading, progress = bytesWritten / contentLength.toFloat(), ) } }, ), ) } }.getOrElse { error -> logger.e(error) return@coroutineScope Result.failure( workDataOf( Parameters.Output.ERROR to "Failed to upload to S3", ), ) } setProgress( uploadId = uploadId, name = uploadName, step = Parameters.Progress.Steps.Finalizing, ) withRetry( retryCount = 5, minBackoffDelay = 2.seconds, ) { uploadsService.finishProcessingUpload( uploadId = uploadDestination.id, finishUploadSpec = FinishUploadSpec( uploadType = "audio_recording", uploadFilename = uploadName, ), ) }.getOrElse { logger.e(it) return@coroutineScope Result.failure( workDataOf( Parameters.Output.ERROR to "Failed to finalize upload", ), ) } var status: UploadRequestStatusSchema? = null fun UploadRequestStatusSchema.isFinishedProcessing() = this.status == "complete" || this.status == "error" while (status?.isFinishedProcessing() != true) { status = withRetry( retryCount = 5, minBackoffDelay = 2.seconds, ) { uploadsService.getUploadStatus(uploadId.value).getOrElse { error -> logger.e(error.toThrowable()) throw error.toThrowable() // caught in withRetry } }.getOrElse { UploadRequestStatusSchema( id = uploadDestination.id, status = "error", ) } if (!status.isFinishedProcessing()) { delay(4.seconds) } } if (status.status == "error") { return@coroutineScope Result.failure( workDataOf( Parameters.Output.ERROR to status.errorMessage, ), ) } val thumbnail = status.imageUrl?.let { Url(it) } setProgress( uploadId = uploadId, name = uploadName, thumbnail = thumbnail, step = Parameters.Progress.Steps.Finalizing, ) val clipId = status.s3Id?.let { Id(it) } ?: return@coroutineScope Result.failure( workDataOf( Parameters.Output.ERROR to "Did not receive s3 id from upload processing", ), ) withRetry( retryCount = 5, minBackoffDelay = 2.seconds, ) { uploadsService.initializeUploadClip( uploadId = uploadId.value, uploadClipSchema = UploadClipSchema( clipId = clipId.value, ), ).getOrElse { error -> logger.e(error.toThrowable()) throw error.toThrowable() // caught in withRetry } }.getOrElse { error -> return@coroutineScope Result.failure( workDataOf( Parameters.Output.ERROR to "Failed to initialize upload clip", ), ) } withRetry( retryCount = 5, minBackoffDelay = 2.seconds, ) { genService.setMetadata( genId = clipId.value, clipMetadataSpec = ClipMetadataSpec( title = uploadName, imageUrl = status.imageUrl, isAudioUploadTosAccepted = true, ), ).getOrElse { error("failed to set metadata") } }.getOrElse { logger.e(it) return@coroutineScope Result.failure( workDataOf( Parameters.Output.ERROR to "Failed to set clip metadata", ), ) } Result.success( workDataOf( Parameters.Output.UPLOAD_ID to uploadId.value, Parameters.Output.CLIP_ID to clipId.value, Parameters.Progress.NAME to uploadName, Parameters.Progress.THUMBNAIL to thumbnail?.url, ), ) } private suspend fun withRetry( retryCount: Int, minBackoffDelay: Duration, lambda: suspend () -> T, ): kotlin.Result { var result: kotlin.Result? = null for (retry in 0 until retryCount) { result = kotlin.runCatching { lambda() } if (result.isSuccess) { break } if (retry < retryCount - 1) { delay(2.0.pow(retry.toDouble()) * minBackoffDelay) } } return result ?: kotlin.Result.failure(IllegalStateException("retryCount must be greater than 0")) } }