// Copyright 2014-2025 Jesper Larsson
//
// This file is part of Klipspringer, <https://klipspringer.avadeaux.net/>
//
// Klipspringer is free software: you can redistribute it and/or modify it under the terms of the
// GNU General Public License as published by the Free Software Foundation, either version 3 of the
// License, or (at your option) any later version.
//
// Klipspringer is distributed in the hope that it will be useful, but WITHOUT ANY WARRANTY; without
// even the implied warranty of MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the GNU
// General Public License for more details.
//
// You should have received a copy of the GNU General Public License along with Klipspringer. If
// not, see <https://www.gnu.org/licenses/>.

#include <limits.h>
#include <stdlib.h>
#include <stdint.h>
#include <string.h>
#include <sys/errno.h>
#include "lame/lame.h"
#include "Library.h"
#include "net_avadeaux_klipspringer_codec_LameStream.h" // generated by javac -h

#if INT_MAX != 0x7fffffff || SHRT_MAX != 0x7fff
#error "Integer sizes do not match assumptions of LAME"
#endif

// A pointer to this struct is cast to an integer and used for interaction with the Java side.
typedef struct {
    lame_global_flags *gfp;                     // Lame encoder record
    bool flush_called;                          // flush has been called
    unsigned mp3buf_size;                       // size of mp3buf
    unsigned char *mp3buf;                      // encoded data sent to channel
    jobject mp3_bb;                             // ByteBuffer wrap of mp3buf, NULL if file not NULL
    jobject receiver;                           // receiver of encoded data, NULL if file not NULL
    jmethodID receiveMid;                       // receciver.receive method, NULL if file not NULL
    FILE *file;                                 // if not NULL, send data to file, not receiver
    unsigned ss;                                // bytes per sample, 2 or 4
    unsigned channels;                          // 1 for mono, 2 for stereo
} EncoderRecord;

// -------------------------------------------------------------------------------------------------
// Helper functions used on errors and by free. The policy is that erec is freed on error in create,
// but in other cases the caller should catch IOException and call free.

// Closes and nulls everything out, to prevent cryptic reuse errors.
static void
free_erec(JNIEnv *env, EncoderRecord *erec) {
    if (erec != NULL) {
        if (!erec->flush_called) {              // finish encoder normally, just in case
            lame_encode_flush(erec->gfp, erec->mp3buf, erec->mp3buf_size);
            erec->flush_called = true;
        }
        if (erec->receiver) { (*env)->DeleteGlobalRef(env, erec->receiver); }
        if (erec->file) { fclose(erec->file); }
        if (erec->gfp) { lame_close(erec->gfp); }
        if (erec->mp3buf) { free(erec->mp3buf); }
        if (erec->mp3_bb) { (*env)->DeleteGlobalRef(env, erec->mp3_bb); }
        free(erec);
    }
}

// Throws Error unless already thrown, frees drec, and returns 0.
static jlong
fail(JNIEnv *env, const char *message, EncoderRecord *erec) {
    raiseError(env, message);
    free_erec(env, erec);
    return 0;
}

// -------------------------------------------------------------------------------------------------
// Native LameStream methods.

// Common part of encodeBuffer and finish, which sends encoded data on to the reciver after handling
// possible encoding errors.
static void
write_mp3buffer(JNIEnv *env, jobject jthis, EncoderRecord *erec, int mp3bytes, jint frames) {
    if (mp3bytes < 0) {
        switch (mp3bytes) {
        case -1:
            raiseErrorCodeException(env, mp3bytes, "MP3 buffer was too small");
            break;
        case -2:
            raiseErrorCodeException(env, mp3bytes, "memory allocation problem in Lame encode");
            break;
        case -3:
            raiseErrorCodeException(env, mp3bytes, "lame_init_params() not called");
            break;
        case -4:
            raiseErrorCodeException(env, mp3bytes, "psycho acoustic problems in Lame encode");
            break;
        default:
            raiseErrorCodeException(env, mp3bytes, "unknown Lame encoding error");
        }
        return;
    }

    if (erec->file == NULL) {
        if (!byteBufferPosition(env, erec->mp3_bb, 0)) { return; }
        if (!byteBufferLimit(env, erec->mp3_bb, mp3bytes)) { return; }
        (*env)->CallVoidMethod(env, erec->receiver, erec->receiveMid, erec->mp3_bb, frames);
    } else {
        size_t n = fwrite(erec->mp3buf, 1, mp3bytes, erec->file);
        if (n < mp3bytes) { raiseErrorCodeException(env, errno, strerror(errno)); }
    }
}

#define METHOD(name) JNICALL Java_net_avadeaux_klipspringer_codec_LameStream_ ## name

JNIEXPORT jlong
METHOD(create) (JNIEnv *env,
                jclass jeclass,
                jbyteArray jfnam,
                jint jrate,
                jint jbips,
                jint jchannels,
                jint jss,
                jboolean jsigned,
                jboolean jbigend,
                jint jbrate,
                jint jquality,
                jint jbufferFrames,
                jobject jreceiver)
{
    if (jchannels != 1 && jchannels != 2) { return fail(env, "Too many channels", NULL); }

    EncoderRecord *erec = malloc(sizeof *erec);
    if (erec == NULL) { return fail(env, "Failed to allocate encoder record", erec); }

    // Null values to make sure free_erec does not segv on error.
    erec->gfp = NULL;
    erec->flush_called = false;
    erec->mp3buf = NULL;
    erec->mp3_bb = NULL;
    erec->receiver = NULL;
    erec->file = NULL;

    // Remember format and channels.
    erec->ss = jss;
    erec->channels = jchannels;

    // Global references to receiver or file.
    if (jfnam == NULL) {
        erec->receiver = (*env)->NewGlobalRef(env, jreceiver);
        if (erec->receiver == NULL) { return fail(env, "Failed to get global reference", erec); }
        erec->receiveMid = (*env)->GetMethodID(env, (*env)->GetObjectClass(env, jreceiver), "receive", "(Ljava/nio/ByteBuffer;I)V");
        if (erec->receiveMid == NULL) { return fail(env, "Failed to get receive method", erec); }
    } else {
        jbyte *fnam = (*env)->GetByteArrayElements(env, jfnam, NULL);
        if (fnam == NULL) { return fail(env, "Failed to allocate filename string", erec); }
        erec->file = fopen((char *) fnam, "w");
        (*env)->ReleaseByteArrayElements(env, jfnam, fnam, JNI_ABORT);
        if (erec->file == NULL) { return fail(env, raiseErrorCodeException(env, errno, strerror(errno)), erec); }
    }

    // Set up encoder.
    if ((erec->gfp = lame_init()) == NULL) {
        return fail(env, "Failed to create Lame encoder", erec);
    }
    if (lame_set_in_samplerate(erec->gfp, jrate) < 0
        || lame_set_num_channels(erec->gfp, 2) < 0
        || lame_set_brate(erec->gfp, jbrate) < 0
        || lame_set_mode(erec->gfp, JOINT_STEREO) < 0
        || lame_set_quality(erec->gfp, jquality) < 0
        || jfnam != NULL && lame_set_bWriteVbrTag(erec->gfp, 0) < 0
        || lame_init_params(erec->gfp) < 0) {
        return fail(env, raiseIOException(env, "Failed to initialize Lame encoder"), erec);
    }

    // Allocate MP3 data buffer for writing to channel.
    erec->mp3buf_size = jbufferFrames + jbufferFrames/4 + 7200;
    erec->mp3buf = malloc(erec->mp3buf_size);
    if (erec->mp3buf == NULL) { return fail(env, "Failed to allocate buffer", erec); }
    if (jfnam == NULL) {
        jobject mp3_bb = wrapByteBuffer(env, erec->mp3buf, erec->mp3buf_size, platform_bigend());
        if (mp3_bb == NULL) { return fail(env, "Failed to wrap in ByteBuffer", erec); }
        erec->mp3_bb = (*env)->NewGlobalRef(env, mp3_bb);
        if (erec->mp3_bb == NULL) { return fail(env, "Failed to wrap in ByteBuffer", erec); }
    }

    return (intptr_t) erec;
}

JNIEXPORT void
METHOD(write) (JNIEnv *env,
               jobject jthis,
               jlong jerec,
               jobject jdata,
               jint jpos,
               jint jframes)
{
    EncoderRecord *erec = (EncoderRecord *) (intptr_t) jerec;

    void *data = (*env)->GetDirectBufferAddress(env, jdata);
    if (data == NULL) { raiseError(env, "Cannot get data byte buffer address"); return; }
    data = ((char *) data) + jpos;

    int n = erec->channels == 2
        ? (erec->ss < 3         // stereo
           ? lame_encode_buffer_interleaved(erec->gfp, data, jframes, erec->mp3buf, erec->mp3buf_size)
           : lame_encode_buffer_interleaved_int(erec->gfp, data, jframes, erec->mp3buf, erec->mp3buf_size))
        : (erec->ss < 3         // mono, use same data for both channels
           ? lame_encode_buffer(erec->gfp, data, data, jframes, erec->mp3buf, erec->mp3buf_size)
           : lame_encode_buffer_int(erec->gfp, data, data, jframes, erec->mp3buf, erec->mp3buf_size));

    write_mp3buffer(env, jthis, erec, n, jframes);
}

JNIEXPORT void
METHOD(finish) (JNIEnv *env,
                jobject jthis,
                jlong jerec)
{
    EncoderRecord *erec = (EncoderRecord *) (intptr_t) jerec;
    if (erec->flush_called) { return; }
    int n = lame_encode_flush(erec->gfp, erec->mp3buf, erec->mp3buf_size);
    erec->flush_called = true;
    write_mp3buffer(env, jthis, erec, n, 0);
    lame_mp3_tags_fid(erec->gfp, erec->file);
    if (erec->file != NULL) {
        fclose(erec->file);
        erec->file = NULL;
    }
}

JNIEXPORT void
METHOD(free) (JNIEnv *env,
              jobject jthis,
              jlong jerec)
{
    free_erec(env, (EncoderRecord *) (intptr_t) jerec);
}
