/*
 * Copyright 2020 The Android Open Source Project
 *
 * Licensed under the Apache License, Version 2.0 (the "License");
 * you may not use this file except in compliance with the License.
 * You may obtain a copy of the License at
 *
 *      http://www.apache.org/licenses/LICENSE-2.0
 *
 * Unless required by applicable law or agreed to in writing, software
 * distributed under the License is distributed on an "AS IS" BASIS,
 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
 * See the License for the specific language governing permissions and
 * limitations under the License.
 */

#include "NativeAudioAnalyzer.h"

// AAudioStream_isMMapUsed() isn't a public symbol, so use dlsym() to access it.
#include <dlfcn.h>
#define LIB_AAUDIO_NAME "libaaudio.so"
#define FUNCTION_IS_MMAP "AAudioStream_isMMapUsed"
static bool (*aaudioStream_isMMap)(AAudioStream *stream) = nullptr;

#include "WavFileCapture.h"

#include <chrono>
#include <memory>
#include <set>

extern WavFileCapture sWavFileCapture;

static const int OUTPUT_CHANNEL_COUNT = 2;

static void convertPcm16ToFloat(const int16_t *source,
                                float *destination,
                                int32_t numSamples) {
    constexpr float scaler = 1.0f / 32768.0f;
    for (int i = 0; i < numSamples; i++) {
        destination[i] = source[i] * scaler;
    }
}

static void convertPcmFloatToPcm16(const float* source,
                                   int16_t* destination,
                                   int32_t numSamples) {
    for (; numSamples > 0; --numSamples) {
        static const float scale = 1 << 15;
        *destination++ = roundf(fmaxf(fminf((*source++) * scale, scale - 1.f), -scale));
    }
}

static void convertPcmFloatToPcm24Packed(const float* source,
                                         uint8_t* destination,
                                         int32_t numSamples) {
    for (; numSamples > 0; --numSamples) {
        static const float scale = 1 << 23;
        int32_t ival = roundf(fmaxf(fminf((*source++) * scale, scale - 1.f), -scale));
#if HAVE_BIG_ENDIAN
        *destination++ = ival >> 16;
        *destination++ = ival >> 8;
        *destination++ = ival;
#else
        *destination++ = ival;
        *destination++ = ival >> 8;
        *destination++ = ival >> 16;
#endif
    }
}

static inline int32_t clamp32_from_float(float f)
{
    static const float scale = (float)(1UL << 31);
    static const float limpos = 1.;
    static const float limneg = -1.;

    if (f <= limneg) {
        return INT32_MIN;
    } else if (f >= limpos) {
        return INT32_MAX;
    }
    f *= scale;
    /* integer conversion is through truncation (though int to float is not).
     * ensure that we round to nearest, ties away from 0.
     */
    return f > 0 ? f + 0.5 : f - 0.5;
}

static void convertPcmFloatToPcm32(const float* source,
                                   int32_t* destination,
                                   int32_t numSamples) {
    for (; numSamples > 0; --numSamples) {
        *destination++ = clamp32_from_float(*source++);
    }
}

// Fill the audio output buffer.
int32_t NativeAudioAnalyzer::readFormattedData(int32_t numFrames) {
    int32_t framesRead = AAUDIO_ERROR_INVALID_FORMAT;
    if (mActualInputFormat == AAUDIO_FORMAT_PCM_I16) {
        framesRead = AAudioStream_read(mInputStream, mInputShortData,
                                       numFrames,
                                       0 /* timeoutNanoseconds */);
    } else if (mActualInputFormat == AAUDIO_FORMAT_PCM_FLOAT) {
        framesRead = AAudioStream_read(mInputStream, mInputFloatData,
                                       numFrames,
                                       0 /* timeoutNanoseconds */);
    } else {
        ALOGE("ERROR actualInputFormat = %d\n", mActualInputFormat);
        assert(false);
    }
    if (framesRead < 0) {
        // Expect INVALID_STATE if STATE_STARTING
        if (mFramesReadTotal > 0) {
            mInputError = framesRead;
            ALOGE("ERROR in read = %d = %s\n", framesRead,
                   AAudio_convertResultToText(framesRead));
        } else {
            framesRead = 0;
        }
    } else {
        mFramesReadTotal += framesRead;
    }
    return framesRead;
}

bool NativeAudioAnalyzer::has24BitSupport(aaudio_format_t format) {
    return (format == AAUDIO_FORMAT_PCM_FLOAT) || (format == AAUDIO_FORMAT_PCM_I24_PACKED)
            || (format == AAUDIO_FORMAT_PCM_I32);
}

aaudio_data_callback_result_t NativeAudioAnalyzer::dataCallbackProc(
        void *audioData,
        int32_t numFrames
) {
    aaudio_data_callback_result_t callbackResult = AAUDIO_CALLBACK_RESULT_CONTINUE;
    const int32_t numSamplesToProcess = numFrames * mActualOutputChannelCount;
    std::unique_ptr<float[]> outputFloatData = std::make_unique<float[]>(numSamplesToProcess);

    // Read audio data from the input stream.
    int32_t actualFramesRead;

    if (numFrames > mInputFramesMaximum) {
        ALOGE("%s() numFrames:%d > mInputFramesMaximum:%d", __func__, numFrames, mInputFramesMaximum);
        mInputError = AAUDIO_ERROR_OUT_OF_RANGE;
        return AAUDIO_CALLBACK_RESULT_STOP;
    }

    if (numFrames > mMaxNumFrames) {
        mMaxNumFrames = numFrames;
    }
    if (numFrames < mMinNumFrames) {
        mMinNumFrames = numFrames;
    }

    // Get atomic snapshot of the relative frame positions so they
    // can be used to calculate timestamp latency.
    int64_t framesRead = AAudioStream_getFramesRead(mInputStream);
    int64_t framesWritten = AAudioStream_getFramesWritten(mOutputStream);
    mWriteReadDelta = framesWritten - framesRead;
    mWriteReadDeltaValid = true;

    // Silence the output.
    int32_t numBytes = numFrames * mActualOutputChannelCount * mBytesPerOutputSample;
    memset(audioData, 0 /* value */, numBytes);

    if (mNumCallbacksToDrain > 0) {
        // Drain the input FIFOs.
        int32_t totalFramesRead = 0;
        do {
            actualFramesRead = readFormattedData(numFrames);
            if (actualFramesRead > 0) {
                totalFramesRead += actualFramesRead;
            } else if (actualFramesRead < 0) {
                callbackResult = AAUDIO_CALLBACK_RESULT_STOP;
            }
            // Ignore errors because input stream may not be started yet.
        } while (actualFramesRead > 0);
        // Only counts if we actually got some data.
        if (totalFramesRead > 0) {
            mNumCallbacksToDrain--;
        }

    } else if (mNumCallbacksToNotRead > 0) {
        // Let the input fill up a bit so we are not so close to the write pointer.
        mNumCallbacksToNotRead--;
    } else if (mNumCallbacksToDiscard > 0) {
        // Ignore. Allow the input to fill back up to equilibrium with the output.
        actualFramesRead = readFormattedData(numFrames);
        if (actualFramesRead < 0) {
            callbackResult = AAUDIO_CALLBACK_RESULT_STOP;
        }
        mNumCallbacksToDiscard--;

    } else {
        // The full duplex stream is now stable so process the audio.
        int32_t numInputBytes = numFrames * mActualInputChannelCount * sizeof(float);
        memset(mInputFloatData, 0 /* value */, numInputBytes);

        int64_t inputFramesWritten = AAudioStream_getFramesWritten(mInputStream);
        int64_t inputFramesRead = AAudioStream_getFramesRead(mInputStream);
        int64_t framesAvailable = inputFramesWritten - inputFramesRead;

        // Read the INPUT data.
        actualFramesRead = readFormattedData(numFrames); // READ
        if (actualFramesRead < 0) {
            callbackResult = AAUDIO_CALLBACK_RESULT_STOP;
        } else {
            if (actualFramesRead < numFrames) {
                if(actualFramesRead < (int32_t) framesAvailable) {
                    ALOGE("insufficient for no reason, numFrames = %d"
                                   ", actualFramesRead = %d"
                                   ", inputFramesWritten = %d"
                                   ", inputFramesRead = %d"
                                   ", available = %d\n",
                           numFrames,
                           actualFramesRead,
                           (int) inputFramesWritten,
                           (int) inputFramesRead,
                           (int) framesAvailable);
                }
                mInsufficientReadCount++;
                mInsufficientReadFrames += numFrames - actualFramesRead; // deficit
                // ALOGE("Error insufficientReadCount = %d\n",(int)mInsufficientReadCount);
            }

            int32_t numSamples = actualFramesRead * mActualInputChannelCount;

            if (mActualInputFormat == AAUDIO_FORMAT_PCM_I16) {
                convertPcm16ToFloat(mInputShortData, mInputFloatData, numSamples);
            }

            // Process the INPUT and generate the OUTPUT.
            mLoopbackProcessor->process(mInputFloatData,
                                               mActualInputChannelCount,
                                               numFrames,
                                               outputFloatData.get(),
                                               mActualOutputChannelCount,
                                               numFrames);

            switch (mActualOutputFormat) {
                case AAUDIO_FORMAT_PCM_I16:
                    convertPcmFloatToPcm16(outputFloatData.get(),
                                           static_cast<int16_t*>(audioData),
                                           numSamplesToProcess);
                    break;
                case AAUDIO_FORMAT_PCM_FLOAT:
                    memcpy(audioData, outputFloatData.get(), numSamplesToProcess * sizeof(float));
                    break;
                case AAUDIO_FORMAT_PCM_I24_PACKED:
                    convertPcmFloatToPcm24Packed(outputFloatData.get(),
                                                 static_cast<uint8_t*>(audioData),
                                                 numSamplesToProcess);
                    break;
                case AAUDIO_FORMAT_PCM_I32:
                    convertPcmFloatToPcm32(outputFloatData.get(),
                                           static_cast<int32_t*>(audioData),
                                           numSamplesToProcess);
                    break;
                default:
                    ALOGE("Unexpected format: %d", mActualOutputFormat);
                    return AAUDIO_CALLBACK_RESULT_STOP;
            }

            sWavFileCapture.captureData(outputFloatData.get(), numFrames);

            mIsDone = mLoopbackProcessor->isDone();
            if (mIsDone) {
                callbackResult = AAUDIO_CALLBACK_RESULT_STOP;
            }
        }
    }
    mFramesWrittenTotal += numFrames;

    return callbackResult;
}

static aaudio_data_callback_result_t s_MyDataCallbackProc(
        AAudioStream * /* outputStream */,
        void *userData,
        void *audioData,
        int32_t numFrames) {
    NativeAudioAnalyzer *myData = (NativeAudioAnalyzer *) userData;
    return myData->dataCallbackProc(audioData, numFrames);
}

static void s_MyErrorCallbackProc(
        AAudioStream * /* stream */,
        void * userData,
        aaudio_result_t error) {
    ALOGE("Error Callback, error: %d\n",(int)error);
    NativeAudioAnalyzer *myData = (NativeAudioAnalyzer *) userData;
    myData->mOutputError = error;
}

bool NativeAudioAnalyzer::isRecordingComplete() {
    return mWhiteNoiseLatencyAnalyzer.isRecordingComplete();
}

int NativeAudioAnalyzer::analyze() {
    mWhiteNoiseLatencyAnalyzer.analyze();
    return getError(); // TODO review
}

double NativeAudioAnalyzer::getLatencyMillis() {
    return mWhiteNoiseLatencyAnalyzer.getMeasuredLatency() * 1000.0 / mOutputSampleRate;
}

double NativeAudioAnalyzer::getConfidence() {
    return mWhiteNoiseLatencyAnalyzer.getMeasuredConfidence();
}

bool NativeAudioAnalyzer::has24BitHardwareSupport() {
    return mHas24BitHardwareSupport;
}

int NativeAudioAnalyzer::getHardwareFormat() {
    return mHardwareFormat;
}

int NativeAudioAnalyzer::getSampleRate() {
    return mOutputSampleRate;
}

aaudio_result_t NativeAudioAnalyzer::openAudio(int inputDeviceId, int outputDeviceId, int format) {
    mInputDeviceId = inputDeviceId;
    mOutputDeviceId = outputDeviceId;

    AAudioStreamBuilder *builder = nullptr;

    mWhiteNoiseLatencyAnalyzer.setup();
    mLoopbackProcessor = &mWhiteNoiseLatencyAnalyzer; // for latency test

    // Use an AAudioStreamBuilder to contain requested parameters.
    aaudio_result_t result = AAudio_createStreamBuilder(&builder);
    if (result != AAUDIO_OK) {
        ALOGE("AAudio_createStreamBuilder() returned %s",
               AAudio_convertResultToText(result));
        return result;
    }

    // Create the OUTPUT stream -----------------------
    AAudioStreamBuilder_setDirection(builder, AAUDIO_DIRECTION_OUTPUT);
    AAudioStreamBuilder_setPerformanceMode(builder, AAUDIO_PERFORMANCE_MODE_LOW_LATENCY);
    AAudioStreamBuilder_setSharingMode(builder, AAUDIO_SHARING_MODE_EXCLUSIVE);
    AAudioStreamBuilder_setFormat(builder, format);
    AAudioStreamBuilder_setChannelCount(builder, OUTPUT_CHANNEL_COUNT); // stereo
    AAudioStreamBuilder_setDataCallback(builder, s_MyDataCallbackProc, this);
    AAudioStreamBuilder_setErrorCallback(builder, s_MyErrorCallbackProc, this);
    AAudioStreamBuilder_setDeviceId(builder, mOutputDeviceId);

    result = AAudioStreamBuilder_openStream(builder, &mOutputStream);
    if (result != AAUDIO_OK) {
        ALOGE("NativeAudioAnalyzer::openAudio() OUTPUT error %s",
               AAudio_convertResultToText(result));
        return result;
    }

    mActualOutputFormat = AAudioStream_getFormat(mOutputStream);
    if (mActualOutputFormat != format) {
        // This must not happen. Audio framework provide format conversion, app should be
        // able to get the requested PCM format. Return earlier if the actual format is different
        // from the requested one.
        ALOGE("The actual output format(%d) is different from the requested one(%d)",
              mActualOutputFormat, format);
        return AAUDIO_ERROR_INTERNAL;
    }
    mActualOutputChannelCount = AAudioStream_getChannelCount(mOutputStream);
    if (mActualOutputChannelCount != OUTPUT_CHANNEL_COUNT) {
        // This must not happen. Return earlier if the actual channel count is different
        // from the requested one.
        ALOGE("The actual output channel count(%d) is different from the requested one(%d)",
              mActualOutputChannelCount, OUTPUT_CHANNEL_COUNT);
        return AAUDIO_ERROR_INTERNAL;
    }
    switch (mActualOutputFormat) {
        case AAUDIO_FORMAT_PCM_I16:
            mBytesPerOutputSample = sizeof(int16_t);
            break;
        case AAUDIO_FORMAT_PCM_FLOAT:
            mBytesPerOutputSample = sizeof(float);
            break;
        case AAUDIO_FORMAT_PCM_I24_PACKED:
            mBytesPerOutputSample = sizeof(uint8_t) * 3;
            break;
        case AAUDIO_FORMAT_PCM_I32:
            mBytesPerOutputSample = sizeof(int32_t);
            break;
        default:
            ALOGE("Unexpected format: %d", mActualOutputFormat);
            return AAUDIO_ERROR_INTERNAL;
    }

    mHardwareFormat = AAudioStream_getHardwareFormat(mOutputStream);
    mHas24BitHardwareSupport = has24BitSupport(mHardwareFormat);

    // Stream Attributes
    mBurstFrames[STREAM_OUTPUT] = AAudioStream_getFramesPerBurst(mOutputStream);
    (void) AAudioStream_setBufferSizeInFrames(mOutputStream,
            mBurstFrames[STREAM_OUTPUT] * kDefaultOutputSizeBursts);
    mCapacityFrames[STREAM_OUTPUT] = AAudioStream_getBufferCapacityInFrames(mOutputStream);
    mIsLowLatencyStream[STREAM_OUTPUT] =
            AAudioStream_getPerformanceMode(mOutputStream) == AAUDIO_PERFORMANCE_MODE_LOW_LATENCY;

    // Late binding for aaudioStream_isMMap()
    if (aaudioStream_isMMap == nullptr) {
        void* libHandle = dlopen(LIB_AAUDIO_NAME, RTLD_NOW);
        aaudioStream_isMMap = (bool (*)(AAudioStream *stream))
                dlsym(libHandle, FUNCTION_IS_MMAP);
        if (aaudioStream_isMMap == nullptr) {
            ALOGE("%s() could not find " FUNCTION_IS_MMAP, __func__);
        }
    }
    mIsMMap[STREAM_OUTPUT] =
            aaudioStream_isMMap != nullptr ? aaudioStream_isMMap(mOutputStream) : false;

    mOutputSampleRate = AAudioStream_getSampleRate(mOutputStream);

    // Create the INPUT stream -----------------------
    AAudioStreamBuilder_setDirection(builder, AAUDIO_DIRECTION_INPUT);
    AAudioStreamBuilder_setFormat(builder, AAUDIO_FORMAT_UNSPECIFIED);
    AAudioStreamBuilder_setSampleRate(builder, mOutputSampleRate); // must match
    AAudioStreamBuilder_setChannelCount(builder, 1); // mono
    AAudioStreamBuilder_setDataCallback(builder, nullptr, nullptr);
    AAudioStreamBuilder_setErrorCallback(builder, nullptr, nullptr);
    AAudioStreamBuilder_setDeviceId(builder, mInputDeviceId);

    result = AAudioStreamBuilder_openStream(builder, &mInputStream);
    if (result != AAUDIO_OK) {
        ALOGE("NativeAudioAnalyzer::openAudio() INPUT error %s",
               AAudio_convertResultToText(result));
        return result;
    }

    // Stream Attributes
    mIsLowLatencyStream[STREAM_INPUT] =
            AAudioStream_getPerformanceMode(mInputStream) == AAUDIO_PERFORMANCE_MODE_LOW_LATENCY;
    mBurstFrames[STREAM_INPUT] = AAudioStream_getFramesPerBurst(mInputStream);
    mCapacityFrames[STREAM_INPUT] = AAudioStream_getBufferCapacityInFrames(mInputStream);
    mIsMMap[STREAM_INPUT] =
            aaudioStream_isMMap != nullptr ? aaudioStream_isMMap(mInputStream) : false;
    int32_t actualCapacity = AAudioStream_getBufferCapacityInFrames(mInputStream);
    (void) AAudioStream_setBufferSizeInFrames(mInputStream, actualCapacity);

    // ------- Setup loopbackData -----------------------------
    mActualInputFormat = AAudioStream_getFormat(mInputStream);
    mActualInputChannelCount = AAudioStream_getChannelCount(mInputStream);

    // Allocate a buffer for the audio data.
    mInputFramesMaximum = 32 * AAudioStream_getFramesPerBurst(mInputStream);

    if (mActualInputFormat == AAUDIO_FORMAT_PCM_I16) {
        mInputShortData = new int16_t[mInputFramesMaximum * mActualInputChannelCount]{};
    }
    mInputFloatData = new float[mInputFramesMaximum * mActualInputChannelCount]{};

    return result;
}

aaudio_result_t NativeAudioAnalyzer::startAudio() {
    mLoopbackProcessor->prepareToTest();

    mWriteReadDeltaValid = false;

    // Start OUTPUT first so INPUT does not overflow.
    aaudio_result_t result = AAudioStream_requestStart(mOutputStream);
    if (result != AAUDIO_OK) {
        stopAudio();
        return result;
    }

    result = AAudioStream_requestStart(mInputStream);
    if (result != AAUDIO_OK) {
        stopAudio();
        return result;
    }

    return result;
}

aaudio_result_t NativeAudioAnalyzer::stopAudio() {
    aaudio_result_t result1 = AAUDIO_OK;
    aaudio_result_t result2 = AAUDIO_OK;
    ALOGD("stopAudio() , minNumFrames = %d, maxNumFrames = %d\n", mMinNumFrames, mMaxNumFrames);
    // Stop OUTPUT first because it uses INPUT.
    if (mOutputStream != nullptr) {
        result1 = AAudioStream_requestStop(mOutputStream);
    }

    // Stop INPUT.
    if (mInputStream != nullptr) {
        result2 = AAudioStream_requestStop(mInputStream);
    }
    return result1 != AAUDIO_OK ? result1 : result2;
}

aaudio_result_t NativeAudioAnalyzer::closeAudio() {
    aaudio_result_t result1 = AAUDIO_OK;
    aaudio_result_t result2 = AAUDIO_OK;
    // Stop and close OUTPUT first because it uses INPUT.
    if (mOutputStream != nullptr) {
        result1 = AAudioStream_close(mOutputStream);
        mOutputStream = nullptr;
    }

    // Stop and close INPUT.
    if (mInputStream != nullptr) {
        result2 = AAudioStream_close(mInputStream);
        mInputStream = nullptr;
    }
    return result1 != AAUDIO_OK ? result1 : result2;
}

// The timestamp latency is the difference between the input
// and output times for a specific frame.
// Start with the position and time from an input timestamp.
// Map the input position to the corresponding position in output
// and calculate its time.
// Use the difference between framesWritten and framesRead to
// convert input positions to output positions.
// Returns -1.0 if the data callback wasn't called or if the stream is closed.
double NativeAudioAnalyzer::measureTimestampLatencyMillis() {
    if (!mWriteReadDeltaValid) return -1.0; // Data callback never called.

    int64_t writeReadDelta = mWriteReadDelta;
    aaudio_result_t result;

    int64_t inputPosition;
    int64_t inputTimeNanos;
    int64_t outputPosition;
    int64_t outputTimeNanos;
    result = AAudioStream_getTimestamp(mInputStream, CLOCK_MONOTONIC, &inputPosition,
                                       &inputTimeNanos);
    if (result != AAUDIO_OK) {
        if (!mIsLowLatencyStream[STREAM_INPUT] || mIsMMap[STREAM_INPUT]) {
            return -1.0; // Stream is closed.
        }
        // The input stream is using low latency mode and it is not mmap stream. It is using
        // fast capture. Fast capture doesn't expose timestamp as the buffer is small so that
        // there is almost no latency. The read position and current timestamp should be good
        // enough to use here.
        inputPosition = AAudioStream_getFramesRead(mInputStream);
        struct timespec ts;
        clock_gettime(CLOCK_MONOTONIC, &ts);
        inputTimeNanos = ts.tv_sec * 1'000'000'000LL + ts.tv_nsec;
    }
    result = AAudioStream_getTimestamp(mOutputStream, CLOCK_MONOTONIC, &outputPosition,
                                       &outputTimeNanos);
    if (result != AAUDIO_OK) {
        return -1.0; // Stream is closed.
    }
    ALOGI("outPos:%jd, outTs:%jd, inPos:%jd inTs:%jd, writeReadDelta:%jd",
          outputPosition, outputTimeNanos, inputPosition, inputTimeNanos, writeReadDelta);

    // Map input frame position to the corresponding output frame.
    int64_t mappedPosition = inputPosition + writeReadDelta;
    // Calculate when that frame will play.
    int32_t sampleRate = getSampleRate();
    int64_t mappedTimeNanos = outputTimeNanos + ((mappedPosition - outputPosition) * 1e9)
            / sampleRate;

    // Latency is the difference in time between when a frame was recorded and
    // when its corresponding echo was played.
    return (mappedTimeNanos - inputTimeNanos) * 1.0e-6; // convert nanos to millis
}