#define LOG_TAG "LemenzoCallAec"

#include "LemenzoCallAec.h"

#include <api/audio/echo_canceller3_factory.h>
#include <cutils/properties.h>
#include <modules/audio_processing/include/audio_processing.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 reserve(size_t count) {
        if (mData.capacity() < count) {
            mData.reserve(count);
        }
    }

    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;
}

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 (!ensureApmLocked()) return;

        if (sampleRate != mRenderRate) {
            mRenderRate = sampleRate;
            mRenderFifo.clear();
        }

        downmixInt16(buffer, frames, channels, mScratchInput);
        mRenderFifo.push(mScratchInput);
        trimFifoLocked(mRenderFifo, sampleRate * 2); // hard cap: 2 seconds
        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 (!ensureApmLocked()) return;

        if (sampleRate != mRenderRate) {
            mRenderRate = sampleRate;
            mRenderFifo.clear();
        }

        downmixFloat(buffer, frames, channels, mScratchInput);
        mRenderFifo.push(mScratchInput);
        trimFifoLocked(mRenderFifo, 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 (!ensureApmLocked()) return false;
        if (!prepareCaptureRateLocked(sampleRate)) return false;

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

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

        // APM requires ~10 ms blocks. If AudioFlinger gives us a smaller first
        // block, hold it until a complete AEC3 frame exists and output silence
        // for only that startup fraction. Thereafter the FIFO provides a fixed
        // sub-10-ms latency instead of leaking unprocessed echo.
        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 (!ensureApmLocked()) return false;
        if (!prepareCaptureRateLocked(sampleRate)) return false;

        downmixFloat(buffer, frames, channels, mScratchInput);
        mCaptureInputFifo.push(mScratchInput);
        trimFifoLocked(mCaptureInputFifo, 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 <= 192000;
    }

    bool ensureApmLocked() {
        if (mApm != nullptr) return true;

        webrtc::AudioProcessing::Config config;
        config.echo_canceller.enabled = true;
        config.echo_canceller.mobile_mode = false; // AEC3, not legacy AECM.
        config.pipeline.multi_channel_render = false;
        config.pipeline.multi_channel_capture = false;
        config.pipeline.capture_downmix_method =
                webrtc::AudioProcessing::Config::Pipeline::DownmixMethod::kAverageChannels;

        mApm = webrtc::AudioProcessingBuilder()
                .SetConfig(config)
                .SetEchoControlFactory(std::make_unique<webrtc::EchoCanceller3Factory>())
                .Create();
        if (mApm == nullptr) {
            ALOGE("failed to create WebRTC AEC3 AudioProcessing instance");
            return false;
        }

        const int init = mApm->Initialize();
        if (init != webrtc::AudioProcessing::kNoError) {
            ALOGE("WebRTC AEC3 Initialize failed: %d", init);
            mApm = nullptr;
            return false;
        }

        mDelayMs = std::clamp(
                property_get_int32("persist.audio.secureline.aec_delay_ms", 0), 0, 500);
        ALOGI("WebRTC AEC3 ready, stream delay hint=%d ms", mDelayMs);
        return true;
    }

    void resetLocked() {
        mApm = nullptr;
        mRenderRate = 0;
        mCaptureRate = 0;
        mRenderFifo.clear();
        mCaptureInputFifo.clear();
        mCaptureOutputFifo.clear();
        mScratchInput.clear();
        mScratchOutput.clear();
        mFrame.clear();
        mLoggedRender = false;
        mLoggedCapture = false;
        (void)ensureApmLocked();
        ALOGI("cellular AEC state reset");
    }

    bool prepareCaptureRateLocked(int sampleRate) {
        if (mCaptureRate == sampleRate) return true;
        mCaptureRate = sampleRate;
        mCaptureInputFifo.clear();
        mCaptureOutputFifo.clear();
        // APM permits parameter changes directly, but Initialize() at a call
        // capture format transition guarantees that stale filter state from an
        // incompatible capture rate cannot contaminate the next stream.
        const int init = mApm->Initialize();
        if (init != webrtc::AudioProcessing::kNoError) {
            ALOGE("WebRTC AEC3 reinitialize failed for capture rate %d: %d", sampleRate, init);
            return false;
        }
        ALOGI("cellular capture configured: %d Hz mono AEC path", sampleRate);
        return true;
    }

    void processRenderLocked() {
        if (mRenderRate <= 0) return;
        const size_t frameSamples = static_cast<size_t>(
                webrtc::AudioProcessing::GetFrameSize(mRenderRate));
        if (frameSamples == 0) return;

        mFrame.resize(frameSamples);
        const webrtc::StreamConfig config(mRenderRate, 1);
        while (mRenderFifo.size() >= frameSamples) {
            if (!mRenderFifo.pop(mFrame.data(), frameSamples)) break;
            const int rc = mApm->ProcessReverseStream(
                    mFrame.data(), config, config, mFrame.data());
            if (rc != webrtc::AudioProcessing::kNoError) {
                ALOGW("ProcessReverseStream failed: %d", rc);
                break;
            }
            if (!mLoggedRender) {
                mLoggedRender = true;
                ALOGI("Call-RX far-end reference active: %d Hz mono, %zu-frame AEC3 blocks",
                        mRenderRate, frameSamples);
            }
        }
    }

    void processCaptureFramesLocked() {
        if (mCaptureRate <= 0) return;
        const size_t frameSamples = static_cast<size_t>(
                webrtc::AudioProcessing::GetFrameSize(mCaptureRate));
        if (frameSamples == 0) return;

        mFrame.resize(frameSamples);
        const webrtc::StreamConfig config(mCaptureRate, 1);
        while (mCaptureInputFifo.size() >= frameSamples) {
            if (!mCaptureInputFifo.pop(mFrame.data(), frameSamples)) break;

            // AEC3 performs adaptive delay estimation internally. Android's
            // AudioFlinger tap is before hardware render and after hardware
            // capture, so 0 is the safest model-independent hint. A property is
            // provided for lab tuning without another OS rebuild if ever needed.
            (void)mApm->set_stream_delay_ms(mDelayMs);
            const int rc = mApm->ProcessStream(
                    mFrame.data(), config, config, mFrame.data());
            if (rc != webrtc::AudioProcessing::kNoError) {
                ALOGW("ProcessStream failed: %d", rc);
                // Preserve intelligibility if AEC3 rejects one block.
                mCaptureOutputFifo.push(mFrame);
                continue;
            }
            mCaptureOutputFifo.push(mFrame);
            if (!mLoggedCapture) {
                mLoggedCapture = true;
                ALOGI("Call-TX AEC3 capture active: %d Hz mono, %zu-frame blocks, delay=%d ms",
                        mCaptureRate, frameSamples, mDelayMs);
            }
        }
        trimFifoLocked(mCaptureOutputFifo, mCaptureRate * 2);
    }

    static void trimFifoLocked(SampleFifo& fifo, size_t maxSamples) {
        if (fifo.size() <= maxSamples) return;
        // A two-second backlog means the call path has been interrupted or its
        // format changed unexpectedly. Dropping stale state is safer than
        // feeding ancient far-end/capture audio to the adaptive filter.
        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;
    rtc::scoped_refptr<webrtc::AudioProcessing> mApm;
    int mRenderRate = 0;
    int mCaptureRate = 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);
}
