#define LOG_TAG "LemenzoCallAec"

#include "LemenzoCallAec.h"

#include <api/audio/echo_canceller3_factory.h>
#include <api/audio/echo_control.h>
#include <cutils/properties.h>
#include <modules/audio_processing/audio_buffer.h>
#include <utils/Log.h>

#include <algorithm>
#include <cmath>
#include <cstddef>
#include <cstdint>
#include <memory>
#include <mutex>
#include <vector>

namespace {

class SampleFifo {
public:
    void clear() {
        mData.clear();
        mRead = 0;
    }

    size_t size() const {
        return mData.size() - mRead;
    }

    void push(const int16_t* data, size_t count) {
        if (data == nullptr || count == 0) return;
        compactIfNeeded(count);
        mData.insert(mData.end(), data, data + count);
    }

    void push(const std::vector<int16_t>& data) {
        push(data.data(), data.size());
    }

    bool pop(int16_t* dest, size_t count) {
        if (dest == nullptr || count > size()) return false;
        std::copy_n(mData.data() + mRead, count, dest);
        mRead += count;
        if (mRead == mData.size()) {
            clear();
        } else if (mRead > 8192 && mRead > mData.size() / 2) {
            compact();
        }
        return true;
    }

    size_t popAvailable(int16_t* dest, size_t count) {
        const size_t actual = std::min(count, size());
        if (actual == 0 || dest == nullptr) return 0;
        std::copy_n(mData.data() + mRead, actual, dest);
        mRead += actual;
        if (mRead == mData.size()) clear();
        return actual;
    }

private:
    void compactIfNeeded(size_t incoming) {
        if (mRead != 0 && (mData.size() + incoming > mData.capacity() || mRead > 8192)) {
            compact();
        }
    }

    void compact() {
        if (mRead == 0) return;
        mData.erase(mData.begin(), mData.begin() + static_cast<ptrdiff_t>(mRead));
        mRead = 0;
    }

    std::vector<int16_t> mData;
    size_t mRead = 0;
};

static int16_t floatToInt16(float sample) {
    const float clamped = std::max(-1.0f, std::min(1.0f, sample));
    const float scaled = clamped * (clamped >= 0.0f ? 32767.0f : 32768.0f);
    return static_cast<int16_t>(std::lrintf(scaled));
}

static float int16ToFloat(int16_t sample) {
    return static_cast<float>(sample) / 32768.0f;
}

static int normalizeAecRate(int rate) {
    if (rate <= 16000) return 16000;
    if (rate <= 32000) return 32000;
    return 48000;
}

class LemenzoCallAec {
public:
    static LemenzoCallAec& instance() {
        static LemenzoCallAec singleton;
        return singleton;
    }

    void reset() {
        std::lock_guard<std::mutex> lock(mMutex);
        resetLocked();
    }

    void renderInt16(const int16_t* buffer, int frames, int channels, int sampleRate) {
        if (!valid(buffer, frames, channels, sampleRate)) return;
        std::lock_guard<std::mutex> lock(mMutex);
        if (!prepareRatesLocked(sampleRate, true)) return;

        downmixInt16(buffer, frames, channels, mScratchInput);
        mRenderFifo.push(mScratchInput);
        trimFifoLocked(mRenderFifo, static_cast<size_t>(sampleRate) * 2);
        processRenderLocked();
    }

    void renderFloat(const float* buffer, int frames, int channels, int sampleRate) {
        if (!valid(buffer, frames, channels, sampleRate)) return;
        std::lock_guard<std::mutex> lock(mMutex);
        if (!prepareRatesLocked(sampleRate, true)) return;

        downmixFloat(buffer, frames, channels, mScratchInput);
        mRenderFifo.push(mScratchInput);
        trimFifoLocked(mRenderFifo, static_cast<size_t>(sampleRate) * 2);
        processRenderLocked();
    }

    bool processCaptureInt16(int16_t* buffer, int frames, int channels, int sampleRate) {
        if (!valid(buffer, frames, channels, sampleRate)) return false;
        std::lock_guard<std::mutex> lock(mMutex);
        if (!prepareRatesLocked(sampleRate, false)) return false;

        downmixInt16(buffer, frames, channels, mScratchInput);
        mCaptureInputFifo.push(mScratchInput);
        trimFifoLocked(mCaptureInputFifo, static_cast<size_t>(sampleRate) * 2);
        processCaptureFramesLocked();

        mScratchOutput.resize(static_cast<size_t>(frames));
        const size_t produced = mCaptureOutputFifo.popAvailable(
                mScratchOutput.data(), static_cast<size_t>(frames));

        std::fill(mScratchOutput.begin() + static_cast<ptrdiff_t>(produced),
                  mScratchOutput.end(), 0);

        for (int f = 0; f < frames; ++f) {
            const int16_t sample = mScratchOutput[static_cast<size_t>(f)];
            for (int c = 0; c < channels; ++c) {
                buffer[static_cast<size_t>(f) * channels + c] = sample;
            }
        }
        return true;
    }

    bool processCaptureFloat(float* buffer, int frames, int channels, int sampleRate) {
        if (!valid(buffer, frames, channels, sampleRate)) return false;
        std::lock_guard<std::mutex> lock(mMutex);
        if (!prepareRatesLocked(sampleRate, false)) return false;

        downmixFloat(buffer, frames, channels, mScratchInput);
        mCaptureInputFifo.push(mScratchInput);
        trimFifoLocked(mCaptureInputFifo, static_cast<size_t>(sampleRate) * 2);
        processCaptureFramesLocked();

        mScratchOutput.resize(static_cast<size_t>(frames));
        const size_t produced = mCaptureOutputFifo.popAvailable(
                mScratchOutput.data(), static_cast<size_t>(frames));
        std::fill(mScratchOutput.begin() + static_cast<ptrdiff_t>(produced),
                  mScratchOutput.end(), 0);

        for (int f = 0; f < frames; ++f) {
            const float sample = int16ToFloat(mScratchOutput[static_cast<size_t>(f)]);
            for (int c = 0; c < channels; ++c) {
                buffer[static_cast<size_t>(f) * channels + c] = sample;
            }
        }
        return true;
    }

private:
    LemenzoCallAec() = default;

    template <typename T>
    static bool valid(const T* buffer, int frames, int channels, int sampleRate) {
        return buffer != nullptr && frames > 0 && channels > 0 && channels <= 8
                && sampleRate >= 8000 && sampleRate <= 48000 && sampleRate % 100 == 0;
    }

    bool prepareRatesLocked(int sampleRate, bool render) {
        if (render) {
            if (mRenderInputRate != sampleRate) {
                mRenderInputRate = sampleRate;
                mRenderFifo.clear();
                mRenderBuffer.reset();
            }
        } else {
            if (mCaptureInputRate != sampleRate) {
                mCaptureInputRate = sampleRate;
                mCaptureInputFifo.clear();
                mCaptureOutputFifo.clear();
                mCaptureBuffer.reset();
            }
        }

        const int desiredRate = std::max(
                mRenderInputRate > 0 ? normalizeAecRate(mRenderInputRate) : 0,
                mCaptureInputRate > 0 ? normalizeAecRate(mCaptureInputRate) : 0);
        if (desiredRate <= 0) return false;

        if (mEcho == nullptr || mProcessingRate != desiredRate) {
            if (!createEchoLocked(desiredRate)) return false;
        }

        if (mRenderInputRate > 0 && mRenderBuffer == nullptr) {
            mRenderBuffer = std::make_unique<webrtc::AudioBuffer>(
                    mRenderInputRate, 1,
                    mProcessingRate, 1,
                    mProcessingRate, 1);
        }

        if (mCaptureInputRate > 0 && mCaptureBuffer == nullptr) {
            mCaptureBuffer = std::make_unique<webrtc::AudioBuffer>(
                    mCaptureInputRate, 1,
                    mProcessingRate, 1,
                    mCaptureInputRate, 1);
        }

        return true;
    }

    bool createEchoLocked(int processingRate) {
        webrtc::EchoCanceller3Factory factory;
        std::unique_ptr<webrtc::EchoControl> echo =
                factory.Create(processingRate, 1, 1);
        if (!echo) {
            ALOGE("failed to create direct WebRTC EchoCanceller3 at %d Hz", processingRate);
            return false;
        }

        mEcho = std::move(echo);
        mEcho->SetCaptureOutputUsage(true);
        mProcessingRate = processingRate;
        // -1 means: no external delay hint.
        // AEC3 then relies on its own adaptive render/capture delay estimator.
        //
        // If a verified platform-specific delay is ever provided through the
        // property, values from 0..500 ms are accepted.
        mDelayMs =
                property_get_int32(
                        "persist.audio.secureline.aec_delay_ms",
                        -1);

        if (mDelayMs < -1 || mDelayMs > 500) {
            ALOGW("invalid external AEC3 delay hint %d ms; using auto delay",
                  mDelayMs);
            mDelayMs = -1;
        }

        mRenderBuffer.reset();
        mCaptureBuffer.reset();
        mRenderFifo.clear();
        mCaptureInputFifo.clear();
        mCaptureOutputFifo.clear();
        mLoggedRender = false;
        mLoggedCapture = false;

        if (mDelayMs >= 0) {
            ALOGI("direct AEC3 ready: %d Hz, external delay hint=%d ms",
                  mProcessingRate, mDelayMs);
        } else {
            ALOGI("direct AEC3 ready: %d Hz, adaptive delay estimation enabled",
                  mProcessingRate);
        }
        return true;
    }

    void resetLocked() {
        mEcho.reset();
        mRenderBuffer.reset();
        mCaptureBuffer.reset();
        mProcessingRate = 0;
        mRenderInputRate = 0;
        mCaptureInputRate = 0;
        mRenderFifo.clear();
        mCaptureInputFifo.clear();
        mCaptureOutputFifo.clear();
        mScratchInput.clear();
        mScratchOutput.clear();
        mFrame.clear();
        mLoggedRender = false;
        mLoggedCapture = false;
        ALOGI("direct cellular AEC3 state reset");
    }

    void processRenderLocked() {
        if (!mEcho || !mRenderBuffer || mRenderInputRate <= 0) return;

        const size_t inputFrameSamples = static_cast<size_t>(mRenderInputRate / 100);
        mFrame.resize(inputFrameSamples);
        const webrtc::StreamConfig inputConfig(mRenderInputRate, 1);

        while (mRenderFifo.size() >= inputFrameSamples) {
            if (!mRenderFifo.pop(mFrame.data(), inputFrameSamples)) break;

            mRenderBuffer->CopyFrom(mFrame.data(), inputConfig);
            if (mRenderBuffer->num_bands() > 1) {
                mRenderBuffer->SplitIntoFrequencyBands();
            }
            mEcho->AnalyzeRender(mRenderBuffer.get());

            if (!mLoggedRender) {
                mLoggedRender = true;
                ALOGI("Call-RX AEC3 reference active: input=%d Hz processing=%d Hz bands=%zu",
                        mRenderInputRate, mProcessingRate, mRenderBuffer->num_bands());
            }
        }
    }

    void processCaptureFramesLocked() {
        if (!mEcho || !mCaptureBuffer || mCaptureInputRate <= 0) return;

        const size_t inputFrameSamples = static_cast<size_t>(mCaptureInputRate / 100);
        mFrame.resize(inputFrameSamples);
        const webrtc::StreamConfig inputConfig(mCaptureInputRate, 1);
        const webrtc::StreamConfig outputConfig(mCaptureInputRate, 1);

        while (mCaptureInputFifo.size() >= inputFrameSamples) {
            if (!mCaptureInputFifo.pop(mFrame.data(), inputFrameSamples)) break;

            mCaptureBuffer->CopyFrom(mFrame.data(), inputConfig);

            // Match WebRTC APM ordering:
            // AnalyzeCapture(full-band) -> split -> ProcessCapture -> merge.
            mEcho->AnalyzeCapture(mCaptureBuffer.get());

            if (mCaptureBuffer->num_bands() > 1) {
                mCaptureBuffer->SplitIntoFrequencyBands();
            }

            // WebRTC defines SetAudioBufferDelay() as an OPTIONAL
            // externally supplied estimate. Do not force 0 ms when the
            // actual Pixel cellular render/capture delay is unknown.
            //
            // Without an external hint, EchoCanceller3 keeps its internal
            // adaptive delay estimator in control.
            if (mDelayMs >= 0) {
                mEcho->SetAudioBufferDelay(mDelayMs);
            }

            mEcho->ProcessCapture(mCaptureBuffer.get(), false);

            if (mCaptureBuffer->num_bands() > 1) {
                mCaptureBuffer->MergeFrequencyBands();
            }

            mCaptureBuffer->CopyTo(outputConfig, mFrame.data());
            mCaptureOutputFifo.push(mFrame);

            if (!mLoggedCapture) {
                mLoggedCapture = true;
                const auto metrics = mEcho->GetMetrics();
                ALOGI("Call-TX direct AEC3 active: input=%d Hz processing=%d Hz bands=%zu "
                      "delayHint=%d ms estimatedDelay=%d ms active=%d",
                        mCaptureInputRate, mProcessingRate, mCaptureBuffer->num_bands(),
                        mDelayMs, metrics.delay_ms, mEcho->ActiveProcessing() ? 1 : 0);
            }
        }

        trimFifoLocked(mCaptureOutputFifo, static_cast<size_t>(mCaptureInputRate) * 2);
    }

    static void trimFifoLocked(SampleFifo& fifo, size_t maxSamples) {
        if (fifo.size() <= maxSamples) return;
        fifo.clear();
        ALOGW("AEC FIFO exceeded safety cap; stale audio discarded");
    }

    static void downmixInt16(
            const int16_t* src, int frames, int channels, std::vector<int16_t>& mono) {
        mono.resize(static_cast<size_t>(frames));
        for (int f = 0; f < frames; ++f) {
            int64_t sum = 0;
            for (int c = 0; c < channels; ++c) {
                sum += src[static_cast<size_t>(f) * channels + c];
            }
            mono[static_cast<size_t>(f)] = static_cast<int16_t>(sum / channels);
        }
    }

    static void downmixFloat(
            const float* src, int frames, int channels, std::vector<int16_t>& mono) {
        mono.resize(static_cast<size_t>(frames));
        for (int f = 0; f < frames; ++f) {
            float sum = 0.0f;
            for (int c = 0; c < channels; ++c) {
                sum += src[static_cast<size_t>(f) * channels + c];
            }
            mono[static_cast<size_t>(f)] = floatToInt16(sum / static_cast<float>(channels));
        }
    }

    std::mutex mMutex;
    std::unique_ptr<webrtc::EchoControl> mEcho;
    std::unique_ptr<webrtc::AudioBuffer> mRenderBuffer;
    std::unique_ptr<webrtc::AudioBuffer> mCaptureBuffer;

    int mProcessingRate = 0;
    int mRenderInputRate = 0;
    int mCaptureInputRate = 0;
    int mDelayMs = 0;

    SampleFifo mRenderFifo;
    SampleFifo mCaptureInputFifo;
    SampleFifo mCaptureOutputFifo;
    std::vector<int16_t> mScratchInput;
    std::vector<int16_t> mScratchOutput;
    std::vector<int16_t> mFrame;
    bool mLoggedRender = false;
    bool mLoggedCapture = false;
};

} // namespace

extern "C" void secureline_call_aec_reset() {
    LemenzoCallAec::instance().reset();
}

extern "C" void secureline_call_aec_render_int16(
        const int16_t* buffer, int frames, int channels, int sampleRate) {
    LemenzoCallAec::instance().renderInt16(buffer, frames, channels, sampleRate);
}

extern "C" void secureline_call_aec_render_float(
        const float* buffer, int frames, int channels, int sampleRate) {
    LemenzoCallAec::instance().renderFloat(buffer, frames, channels, sampleRate);
}

extern "C" bool secureline_call_aec_process_capture_int16(
        int16_t* buffer, int frames, int channels, int sampleRate) {
    return LemenzoCallAec::instance().processCaptureInt16(buffer, frames, channels, sampleRate);
}

extern "C" bool secureline_call_aec_process_capture_float(
        float* buffer, int frames, int channels, int sampleRate) {
    return LemenzoCallAec::instance().processCaptureFloat(buffer, frames, channels, sampleRate);
}
