#include "SecureLineVoiceDSP.h"

#include <cstring>

SecureLineVoiceDSP::SecureLineVoiceDSP()
{
    resetState();
}

float SecureLineVoiceDSP::clampf(
        float v,
        float lo,
        float hi)
{
    return std::max(lo, std::min(hi, v));
}

float SecureLineVoiceDSP::cubic(
        float y0,
        float y1,
        float y2,
        float y3,
        float mu)
{
    const float a0 = y3 - y2 - y0 + y1;
    const float a1 = y0 - y1 - a0;
    const float a2 = y2 - y0;
    const float a3 = y1;
    const float mu2 = mu * mu;

    return
        (a0 * mu * mu2) +
        (a1 * mu2) +
        (a2 * mu) +
        a3;
}

void SecureLineVoiceDSP::setPitch(float p)
{
    mPitch = clampf(p, 0.70f, 1.80f);
}

void SecureLineVoiceDSP::setFormant(float f)
{
    mFormant = clampf(f, 0.65f, 1.25f);
    mFilterConfigDirty = true;
}

void SecureLineVoiceDSP::setVariation(float v)
{
    mVariation = clampf(v, 0.0f, 0.010f);
}

void SecureLineVoiceDSP::setWarmth(float v)
{
    mWarmth = clampf(v, 0.0f, 0.50f);
}

void SecureLineVoiceDSP::setCompression(float v)
{
    mCompression = clampf(v, 0.0f, 1.0f);
}

void SecureLineVoiceDSP::setDrive(float v)
{
    mDrive = clampf(v, 0.0f, 0.35f);
}

void SecureLineVoiceDSP::setOutputGain(float v)
{
    mOutputGain = clampf(v, 0.50f, 1.25f);
}

void SecureLineVoiceDSP::setSpectralTilt(float v)
{
    mSpectralTilt = clampf(v, -1.0f, 1.0f);
}

void SecureLineVoiceDSP::setPresence(float v)
{
    mPresence = clampf(v, -1.0f, 1.0f);
}

void SecureLineVoiceDSP::setSpeakerRoute(bool enabled)
{
    mSpeakerRoute = enabled;
}

void SecureLineVoiceDSP::resetState()
{
    for (int c = 0; c < MAX_CHANNELS; ++c)
    {
        ChannelState& s = mChannels[c];

        std::fill_n(
            s.ring,
            RING_SIZE,
            0.0f
        );

        s.writePos = 0;
        s.grainPhase = 0.0f;
        s.lfoPhase = 0.0f;
        s.hpPrevIn = 0.0f;
        s.hpPrevOut = 0.0f;
        s.warmthLp = 0.0f;
        s.airLp = 0.0f;
        s.compEnv = 0.0f;
        s.tiltLow = 0.0f;
        s.presenceLow = 0.0f;
        s.formant1.reset();
        s.formant2.reset();
        s.formant3.reset();
    }
}

void SecureLineVoiceDSP::configurePeaking(
        Biquad& b,
        float sampleRate,
        float frequency,
        float q,
        float gainDb)
{
    frequency = clampf(
        frequency,
        80.0f,
        sampleRate * 0.45f
    );

    q = clampf(q, 0.25f, 8.0f);

    const float A =
        std::pow(10.0f, gainDb / 40.0f);

    const float w0 =
        2.0f * PI * frequency / sampleRate;

    const float alpha =
        std::sin(w0) / (2.0f * q);

    const float cosw0 = std::cos(w0);

    float b0 = 1.0f + alpha * A;
    float b1 = -2.0f * cosw0;
    float b2 = 1.0f - alpha * A;
    float a0 = 1.0f + alpha / A;
    float a1 = -2.0f * cosw0;
    float a2 = 1.0f - alpha / A;

    const float invA0 = 1.0f / a0;

    b.b0 = b0 * invA0;
    b.b1 = b1 * invA0;
    b.b2 = b2 * invA0;
    b.a1 = a1 * invA0;
    b.a2 = a2 * invA0;
}

void SecureLineVoiceDSP::configureIfNeeded(
        int sampleRate,
        int channels)
{
    sampleRate =
        std::max(8000, std::min(192000, sampleRate));

    channels =
        std::max(1, std::min(MAX_CHANNELS, channels));

    if (sampleRate != mConfiguredSampleRate ||
        channels != mConfiguredChannels)
    {
        mConfiguredSampleRate = sampleRate;
        mConfiguredChannels = channels;
        resetState();
        mFilterConfigDirty = true;
    }

    if (!mFilterConfigDirty)
        return;

    /*
     * Vocal-formant shaping.
     *
     * This is deliberately gentler than a "robot voice" filter:
     * three broad resonant areas follow the formant factor and
     * reshape the vocal envelope after pitch transposition.
     */
    const float formant =
        clampf(mFormant, 0.65f, 1.25f);

    const float f1 = 520.0f * formant;
    const float f2 = 1450.0f * formant;
    const float f3 = 2450.0f * formant;

    const float depth =
        clampf((1.0f - formant) * 10.0f, -2.0f, 3.5f);

    for (int c = 0; c < channels; ++c)
    {
        configurePeaking(
            mChannels[c].formant1,
            (float)sampleRate,
            f1,
            0.85f,
            1.0f + depth
        );

        configurePeaking(
            mChannels[c].formant2,
            (float)sampleRate,
            f2,
            1.00f,
            0.7f + depth * 0.75f
        );

        configurePeaking(
            mChannels[c].formant3,
            (float)sampleRate,
            f3,
            1.15f,
            0.4f + depth * 0.50f
        );
    }

    mFilterConfigDirty = false;
}

float SecureLineVoiceDSP::readRing(
        const ChannelState& state,
        float position) const
{
    while (position < 0.0f)
        position += (float)RING_SIZE;

    while (position >= (float)RING_SIZE)
        position -= (float)RING_SIZE;

    const int i1 = (int)position;
    const int i0 =
        (i1 - 1 + RING_SIZE) % RING_SIZE;
    const int i2 =
        (i1 + 1) % RING_SIZE;
    const int i3 =
        (i1 + 2) % RING_SIZE;

    const float frac =
        position - (float)i1;

    return cubic(
        state.ring[i0],
        state.ring[i1],
        state.ring[i2],
        state.ring[i3],
        frac
    );
}

float SecureLineVoiceDSP::processOne(
        ChannelState& s,
        float input,
        int sampleRate)
{
    input = clampf(input, -1.0f, 1.0f);

    /*
     * First remove rumble/DC before the pitch stage.
     * ~65 Hz high-pass at normal Android capture rates.
     */
    const float hpR =
        std::exp(
            -2.0f * PI * 65.0f /
            (float)sampleRate
        );

    float clean =
        input -
        s.hpPrevIn +
        hpR * s.hpPrevOut;

    s.hpPrevIn = input;
    s.hpPrevOut = clean;

    /*
     * Write once into the delay/grain reservoir.
     */
    s.ring[s.writePos] = clean;

    /*
     * Service pitch is a "depth factor":
     * > 1.0 means a deeper output voice.
     *
     * Convert to conventional output/input pitch ratio.
     */
    float ratio =
        1.0f / clampf(
            mPitch,
            0.70f,
            1.80f
        );

    /*
     * Very small, slow modulation makes the stronger privacy
     * profiles less static while remaining natural.
     */
    const float lfoHz = 0.42f;

    s.lfoPhase +=
        2.0f * PI * lfoHz /
        (float)sampleRate;

    if (s.lfoPhase >= 2.0f * PI)
        s.lfoPhase -= 2.0f * PI;

    ratio *=
        1.0f +
        mVariation *
        std::sin(s.lfoPhase);

    ratio = clampf(ratio, 0.52f, 1.35f);

    /*
     * Dual-grain delay pitch shifter.
     * Two 50%-offset Hann windows crossfade so the stream
     * remains continuous and does not change speaking speed.
     */
    const float minDelay =
        std::max(
            48.0f,
            0.008f * sampleRate
        );

    const float grainRange =
        std::min(
            (float)RING_SIZE * 0.42f,
            std::max(
                384.0f,
                0.050f * sampleRate
            )
        );

    const float delta =
        std::fabs(1.0f - ratio);

    float shifted = clean;

    if (delta > 0.0025f)
    {
        float phaseA = s.grainPhase;
        float phaseB = phaseA + 0.5f;

        if (phaseB >= 1.0f)
            phaseB -= 1.0f;

        float delayA;
        float delayB;

        if (ratio < 1.0f)
        {
            delayA =
                minDelay +
                phaseA * grainRange;

            delayB =
                minDelay +
                phaseB * grainRange;
        }
        else
        {
            delayA =
                minDelay +
                (1.0f - phaseA) * grainRange;

            delayB =
                minDelay +
                (1.0f - phaseB) * grainRange;
        }

        const float readA =
            (float)s.writePos - delayA;

        const float readB =
            (float)s.writePos - delayB;

        const float a =
            readRing(s, readA);

        const float b =
            readRing(s, readB);

        const float wA =
            std::sin(PI * phaseA);

        const float wB =
            std::sin(PI * phaseB);

        const float winA = wA * wA;
        const float winB = wB * wB;

        shifted =
            a * winA +
            b * winB;

        s.grainPhase +=
            delta / grainRange;

        while (s.grainPhase >= 1.0f)
            s.grainPhase -= 1.0f;
    }

    s.writePos =
        (s.writePos + 1) %
        RING_SIZE;

    /*
     * Vocal-envelope / formant character.
     */
    float voiced =
        s.formant1.process(shifted);

    voiced =
        s.formant2.process(voiced);

    voiced =
        s.formant3.process(voiced);

    /*
     * Warmth: subtle low-frequency body, not a muddy bass boost.
     */
    const float warmthAlpha =
        std::exp(
            -2.0f * PI * 240.0f /
            (float)sampleRate
        );

    s.warmthLp =
        warmthAlpha * s.warmthLp +
        (1.0f - warmthAlpha) * voiced;

    voiced +=
        mWarmth *
        0.55f *
        s.warmthLp;

    /*
     * Gentle high-frequency smoothing on the darkest profiles.
     */
    const float airAlpha =
        std::exp(
            -2.0f * PI * 6200.0f /
            (float)sampleRate
        );

    s.airLp =
        airAlpha * s.airLp +
        (1.0f - airAlpha) * voiced;

    const float darkAmount =
        clampf(
            (1.0f - mFormant) * 0.28f,
            0.0f,
            0.12f
        );

    voiced =
        voiced * (1.0f - darkAmount) +
        s.airLp * darkAmount;

    /*
     * Spectral tilt.
     *
     * Negative values move energy toward the low / low-mid region.
     * Positive values reduce low-mid dominance and open the top end.
     * This helps make the five profiles feel like genuinely different
     * vocal characters rather than five pitch presets.
     */
    const float tiltAlpha =
        std::exp(
            -2.0f * PI * 1200.0f /
            (float)sampleRate
        );

    s.tiltLow =
        tiltAlpha * s.tiltLow +
        (1.0f - tiltAlpha) * voiced;

    const float tiltHigh =
        voiced - s.tiltLow;

    if (mSpectralTilt < 0.0f)
    {
        const float amount =
            -mSpectralTilt;

        voiced =
            voiced +
            amount * 0.34f * s.tiltLow -
            amount * 0.18f * tiltHigh;
    }
    else if (mSpectralTilt > 0.0f)
    {
        const float amount =
            mSpectralTilt;

        voiced =
            voiced -
            amount * 0.24f * s.tiltLow +
            amount * 0.28f * tiltHigh;
    }

    /*
     * Speech-presence shaping around the upper-mid speech region.
     * Kept deliberately subtle to preserve intelligibility and avoid
     * harsh "telephone" character.
     */
    const float presenceAlpha =
        std::exp(
            -2.0f * PI * 3200.0f /
            (float)sampleRate
        );

    s.presenceLow =
        presenceAlpha * s.presenceLow +
        (1.0f - presenceAlpha) * voiced;

    const float presenceBand =
        voiced - s.presenceLow;

    voiced +=
        mPresence *
        0.16f *
        presenceBand;

    /*
     * Speech compressor.
     */
    const float absx = std::fabs(voiced);

    const float attack =
        std::exp(
            -1.0f /
            (0.006f * sampleRate)
        );

    const float release =
        std::exp(
            -1.0f /
            (0.090f * sampleRate)
        );

    const float envCoeff =
        (absx > s.compEnv)
        ? attack
        : release;

    s.compEnv =
        envCoeff * s.compEnv +
        (1.0f - envCoeff) * absx;

    const float threshold =
        0.16f;

    const float effectiveCompression =
        mCompression;

    float compGain = 1.0f;

    if (s.compEnv > threshold)
    {
        const float ratioComp =
            1.0f +
            4.0f * effectiveCompression;

        const float target =
            threshold +
            (s.compEnv - threshold) /
            ratioComp;

        compGain =
            target /
            std::max(s.compEnv, 1.0e-6f);
    }

    voiced *=
        1.0f -
        effectiveCompression * 0.18f +
        effectiveCompression * 0.18f * compGain;

    voiced *= compGain;

    /*
     * Mild harmonic density / "studio" body.
     */
    const float effectiveDrive =
        mDrive;

    const float drive =
        1.0f + effectiveDrive * 3.0f;

    voiced =
        std::tanh(voiced * drive) /
        std::max(
            1.0f,
            std::tanh(drive)
        );

    /*
     * Output makeup + soft limiter.
     */
    const float effectiveOutputGain =
        mOutputGain;

    voiced *= effectiveOutputGain;

    const float limited =
        std::tanh(voiced * 1.10f) /
        std::tanh(1.10f);

    return clampf(
        limited * 0.97f,
        -0.985f,
        0.985f
    );
}

void SecureLineVoiceDSP::processFloat(
        float* buffer,
        int frames,
        int channels,
        int sampleRate)
{
    if (!buffer ||
        frames <= 0 ||
        channels <= 0 ||
        channels > MAX_CHANNELS)
    {
        return;
    }

    configureIfNeeded(
        sampleRate,
        channels
    );

    for (int frame = 0;
         frame < frames;
         ++frame)
    {
        for (int channel = 0;
             channel < channels;
             ++channel)
        {
            const int index =
                frame * channels +
                channel;

            buffer[index] =
                processOne(
                    mChannels[channel],
                    buffer[index],
                    sampleRate
                );
        }
    }
}

void SecureLineVoiceDSP::processInt16(
        int16_t* buffer,
        int frames,
        int channels,
        int sampleRate)
{
    if (!buffer ||
        frames <= 0 ||
        channels <= 0 ||
        channels > MAX_CHANNELS)
    {
        return;
    }

    configureIfNeeded(
        sampleRate,
        channels
    );

    for (int frame = 0;
         frame < frames;
         ++frame)
    {
        for (int channel = 0;
             channel < channels;
             ++channel)
        {
            const int index =
                frame * channels +
                channel;

            const float input =
                (float)buffer[index] /
                32768.0f;

            const float output =
                processOne(
                    mChannels[channel],
                    input,
                    sampleRate
                );

            const float scaled =
                clampf(
                    output * 32767.0f,
                    -32768.0f,
                    32767.0f
                );

            buffer[index] =
                (int16_t)scaled;
        }
    }
}
