package com.suno.android.common_data.media.usecase import android.content.Context import android.media.MediaCodec import android.media.MediaCodecList import android.media.MediaExtractor import android.media.MediaFormat import android.net.Uri import androidx.core.net.toUri import dagger.Lazy import dagger.hilt.android.qualifiers.ApplicationContext import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.withContext import java.io.DataOutput import java.io.File import java.io.RandomAccessFile import java.nio.ByteBuffer import java.nio.ByteOrder import javax.inject.Inject import kotlin.time.Duration import kotlin.time.Duration.Companion.microseconds class TrimAudioUseCase @Inject constructor( @ApplicationContext private val context: Context, private val trimWavUseCase: Lazy, private val getAudioTrackIndex: GetAudioTrackIndexUseCase, ) { suspend operator fun invoke( inputFile: Uri, trim: ClosedRange, outputFile: File, ) = withContext(Dispatchers.IO) { if (isWavFile(inputFile)) { trimWavUseCase.get().invoke(inputFile, outputFile.toUri(), trim) } else { trimGeneric(inputFile, trim, outputFile) } } private suspend fun trimGeneric( uri: Uri, trim: ClosedRange, outputFile: File, ) = withContext(Dispatchers.IO) { runCatching { val mediaExtractor = MediaExtractor() mediaExtractor.setDataSource(context, uri, emptyMap()) val index = getAudioTrackIndex(mediaExtractor) ?: error("could not find audio track") mediaExtractor.selectTrack(index) // setup MediaCodec val format = mediaExtractor.getTrackFormat(index) val sampleRate = format.getInteger(MediaFormat.KEY_SAMPLE_RATE) val channelCount = format.getInteger(MediaFormat.KEY_CHANNEL_COUNT) val sampleSize = 16 // pretty sure this is guaranteed by MediaCodec val codec = MediaCodecList(MediaCodecList.ALL_CODECS).findDecoderForFormat(format) val mediaCodec = MediaCodec.createByCodecName(codec) mediaCodec.configure(format, null, null, 0) mediaCodec.start() // skip to start time while (mediaExtractor.sampleTime.microseconds < trim.start) { mediaExtractor.advance() } // setup output file val fileOutputStream = RandomAccessFile(outputFile, "rw") // skip first 44 bytes to make room for wav header (written at end once we know the length) fileOutputStream.skipBytes(44) val byteArray = ByteArray(1024) // write MediaExtractor output to MediaCodec input, MediaCodecOutput to File var hasInputData = true var hasOutputData = true var totalLength = 0 val bufferInfo = MediaCodec.BufferInfo() while (hasInputData || hasOutputData) { if (hasInputData) { val inBufferIndex = mediaCodec.dequeueInputBuffer(100) if (inBufferIndex >= 0) { mediaCodec.getInputBuffer(inBufferIndex)?.let { codecInBuffer -> val len = mediaExtractor.readSampleData(codecInBuffer, 0) // check if this is/should be the end of input and signal end of buffer val hasMoreInput = mediaExtractor.advance() val flags = if (!hasMoreInput || mediaExtractor.sampleTime.microseconds > trim.endInclusive) { hasInputData = false MediaCodec.BUFFER_FLAG_END_OF_STREAM } else { 0 } mediaCodec.queueInputBuffer(inBufferIndex, 0, len, mediaExtractor.sampleTime, flags) } } } if (hasOutputData) { val outBufferIndex = mediaCodec.dequeueOutputBuffer(bufferInfo, 100) if (outBufferIndex >= 0) { mediaCodec.getOutputBuffer(outBufferIndex)?.let { codecOutBuffer -> codecOutBuffer.order(ByteOrder.LITTLE_ENDIAN) var amountRead = 0 while (amountRead < bufferInfo.size) { val amountToRead = (bufferInfo.size - amountRead).coerceAtMost(byteArray.size) // raw pcm codecOutBuffer.get(byteArray, 0, amountToRead) amountRead += amountToRead fileOutputStream.write(byteArray, 0, amountToRead) } totalLength += bufferInfo.size } mediaCodec.releaseOutputBuffer(outBufferIndex, false) } // when output signals end of stream, update state to stop reading outbut data hasOutputData = (bufferInfo.flags and MediaCodec.BUFFER_FLAG_END_OF_STREAM) == 0 } } mediaCodec.release() mediaExtractor.release() fileOutputStream.seek(0) writeWavHeader( fileOutputStream, dataLength = totalLength, sampleRate = sampleRate, sampleSize = sampleSize, channels = channelCount, ) fileOutputStream.close() } } private fun isWavFile( uri: Uri, ): Boolean { context.contentResolver.openInputStream(uri)?.use { input -> val output = ByteArray(12) input.read(output, 0, output.size) return output.toString(Charsets.UTF_8).run { startsWith("RIFF") && endsWith("WAVE") } } return false } private fun writeWavHeader( output: DataOutput, dataLength: Int, sampleRate: Int, sampleSize: Int, channels: Int, ) { val totalSize = 36 + dataLength val byteRate = sampleRate * channels * sampleSize / 8 val header = ByteBuffer.allocate(44).order(ByteOrder.LITTLE_ENDIAN).apply { put("RIFF".toByteArray()) putInt(totalSize) put("WAVE".toByteArray()) put("fmt ".toByteArray()) putInt(16) // Subchunk1Size for PCM putShort(1) // AudioFormat (1 = PCM) putShort(channels.toShort()) putInt(sampleRate) putInt(byteRate) putShort((channels * sampleSize / 8).toShort()) putShort(sampleSize.toShort()) put("data".toByteArray()) putInt(dataLength) }.array() output.write(header) } }