// 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 <time.h>
#include <stdlib.h>
#include <stdint.h>
#include <string.h>
#include <sys/errno.h>
#include "vorbis/vorbisenc.h"
#include "vorbis_error.h"
#include "Library.h"
#include "net_avadeaux_klipspringer_codec_VorbisStream.h" // generated by javac -h

// Polymorphic function types.
struct EncoderRecord;
typedef void (*write_fun) (JNIEnv *, struct EncoderRecord *, const void *restrict, unsigned, unsigned);
typedef void (*transfer_fun) (float **, void *, unsigned, unsigned);

// A pointer to this struct is cast to an integer and used for interaction with the Java side.
typedef struct EncoderRecord {
    vorbis_info vi;                             // Vorbis encoding info
    vorbis_dsp_state vd;                        // Vorbis encoding state
    vorbis_block vb;                            // Vorbis block record reused during encoding
    ogg_stream_state os;                        // state for Ogg output stream
    ogg_int64_t last_pos;                       // number of sample frames written so far
    bool finish_called;                         // finish has been called
    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 channels;                          // number of channels, 2 for stereo
    write_fun write;                            // polymorphic write function
    transfer_fun transfer;                      // polymorphic sample transfer function
} 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.

static void finish(JNIEnv *, EncoderRecord *);
static void write_ignore(JNIEnv *, EncoderRecord *, const void *restrict, unsigned, unsigned);

// Closes and nulls everything out, to prevent cryptic reuse errors.
static void
free_erec(JNIEnv *env, EncoderRecord *erec) {
    if (erec != NULL) {
        erec-> write = write_ignore;
        finish(env, erec);                      // finish encoder normally, just in case
        if (erec->receiver) { (*env)->DeleteGlobalRef(env, erec->receiver); }
        if (erec->file) { fclose(erec->file); }
        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;
}

// -------------------------------------------------------------------------------------------------
// Buffer write methods.

// Version of erec->write used when encoding to a file.
static void
write_file(JNIEnv *env, EncoderRecord *erec, const void *restrict p, unsigned bytes, unsigned frames) {
    errno = 0;
    if (fwrite(p, 1, bytes, erec->file) < bytes) {
        raiseIOException(env, errno ? strerror(errno) : "write truncated");
    }
}

// Version of erec->write used when encoding to a
// net.avadeaux.klipspringer.codec.ExternalEncodedStream.Receiver.
static void
write_receiver(JNIEnv *env, EncoderRecord *erec, const void *restrict p, unsigned bytes, unsigned frames) {
    jobject bb = wrapByteBuffer(env, (void *restrict) p, bytes, false);
    if (bb) { (*env)->CallVoidMethod(env, erec->receiver, erec->receiveMid, bb, frames); }
}

// Version of erec->write used during free, when there should be no actual write to output.
static void
write_ignore(JNIEnv *env, EncoderRecord *erec, const void *restrict p, unsigned bytes, unsigned frames) {
    // ignore
}

// -------------------------------------------------------------------------------------------------
// Sample transfer methods for different sample sizes and number of channels. Instantiations for
// fixed number of channels are not necessary, but allow the compiler to produce better code. See:
// https://klipspringer.avadeaux.net/optimization-by-doing-the-same-thing-and-expecting-different-results/

#define TRANSFER(b, c) transfer ## _ ## b ## _ ## c

#define TRANSFER_FUN_BC(b, s, c)                                        \
    static void                                                         \
    TRANSFER(b, c) (float **buf, void *data, unsigned channels, unsigned n) {     \
        int ## b ## _t *pcmbuf = data;                                  \
        for (unsigned i = 0; i < n; i++) {                                   \
            for (unsigned j = 0; j < c; j++) {                               \
                buf[j][i] = pcmbuf[channels*i + j]/s;                   \
            }                                                           \
        }                                                               \
    }

#define TRANSFER_FUN_C(c)                                               \
    TRANSFER_FUN_BC(8, 256.0f, c)                                       \
    TRANSFER_FUN_BC(16, 32768.0f, c)                                    \
    TRANSFER_FUN_BC(32, 2147483648.0f, c)

TRANSFER_FUN_C(1)
TRANSFER_FUN_C(2)
TRANSFER_FUN_C(channels)

// -------------------------------------------------------------------------------------------------
// Native VorbisStream methods.

#define METHOD(name) JNICALL Java_net_avadeaux_klipspringer_codec_VorbisStream_ ## name

// Macro that sets erec->transfer to the correct value for the supplied sample size
#define CHOOSE_TRANSFER(c)                                              \
    switch (jss) {                                                      \
    case 1: erec->transfer = TRANSFER(8, c); break;                     \
    case 2: erec->transfer = TRANSFER(16, c); break;                    \
    case 4: erec->transfer = TRANSFER(32, c); break;                    \
    default: return fail(env, "invalid sample size", erec);             \
    }

JNIEXPORT jlong
METHOD(create) (JNIEnv *env,
                jclass jeclass,
                jbyteArray jfnam,
                jint jrate,
                jint jbips,
                jint jchannels,
                jint jss,
                jfloat jquality,
                jobject jreceiver)
{
    EncoderRecord *erec = malloc(sizeof *erec);
    if (erec == NULL) { return fail(env, "Failed to allocate encoder record", erec); }

    // Set values to make sure free_erec does not segv on error.
    erec->finish_called = true;
    erec->receiver = NULL;
    erec->file = NULL;

    // Remember format and channels.
    erec->channels = jchannels;
    switch (jchannels) {
    case 1: CHOOSE_TRANSFER(1); break;
    case 2: CHOOSE_TRANSFER(2); break;
    default: CHOOSE_TRANSFER(channels); break;
    }

    // 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); }
        erec->write = write_receiver;
    } 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); }
        erec->write = write_file;
    }

    // Initialize vorbis records.
    vorbis_info_init(&erec->vi);
    srand(time(NULL));
    int err = vorbis_encode_init_vbr(&erec->vi, jchannels, jrate, jquality);
    if (err < 0
        || vorbis_analysis_init(&erec->vd, &erec->vi)
        || vorbis_block_init(&erec->vd, &erec->vb)
        || ogg_stream_init(&erec->os, rand())) {
        return fail(env, vorbis_strerror(err), erec);
    }

    // Output header.
    char *comment = "klipspringer.avadeaux.net";
    vorbis_comment vc = { &comment, (int[]) { (int) strlen(comment) }, 1, NULL };
    ogg_packet op, op_comm, op_code;
    err = 0;
    if ((err = vorbis_analysis_headerout(&erec->vd, &vc, &op, &op_comm, &op_code) < 0)
        || ogg_stream_packetin(&erec->os, &op)
        || ogg_stream_packetin(&erec->os, &op_comm)
        || ogg_stream_packetin(&erec->os, &op_code)) {
        return fail(env, vorbis_strerror(err), erec);
    }

    // Clear page.
    ogg_page og;
    while (ogg_stream_flush(&erec->os, &og)) {
        erec->write(env, erec, og.header, og.header_len, 0);
        erec->write(env, erec, og.body, og.body_len, 0);
    }

    erec->finish_called = false;
    return (intptr_t) erec;
}

// Subroutine for write and finish.
static bool
analysis_write(JNIEnv *env, EncoderRecord *erec, void *restrict data, unsigned n) {
    float **buf = vorbis_analysis_buffer(&erec->vd, n);
    erec->transfer(buf, data, erec->channels, n);
    int err = vorbis_analysis_wrote(&erec->vd, n);
    if (err < 0) {
        raiseErrorCodeException(env, err, vorbis_strerror(err));
        return false;
    }
    return true;
}

// Subroutine for write and finish.
static void
write_block(JNIEnv *env, EncoderRecord *erec) {
    ogg_packet op;
    ogg_page og;
    int err;

    for (int blk; (blk = vorbis_analysis_blockout(&erec->vd, &erec->vb));) {
        if (blk < 0) { raiseErrorCodeException(env, blk, vorbis_strerror(blk)); return; }
        if ((err = vorbis_analysis(&erec->vb, NULL)) < 0
            || (err = vorbis_bitrate_addblock(&erec->vb)) < 0) { raiseErrorCodeException(env, err, vorbis_strerror(err)); return; }
        for (int pkt; (pkt = vorbis_bitrate_flushpacket(&erec->vd, &op));) {
            if (pkt == 0) { break; }
            if (pkt < 0) { raiseErrorCodeException(env, blk, vorbis_strerror(pkt)); return; }
            if (ogg_stream_packetin(&erec->os, &op) < 0) { raiseIOException(env, "error in Ogg stream encoding"); return; }
            while (ogg_stream_pageout(&erec->os, &og)) {
                unsigned frames = 0;
                ogg_int64_t pos = ogg_page_granulepos(&og);
                if (pos > -1) {
                    frames = (unsigned) (pos - erec->last_pos);
                    erec->last_pos = pos;
                }
                erec->write(env, erec, og.header, og.header_len, 0);
                erec->write(env, erec, og.body, og.body_len, frames);
            }
        }
    }
}

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"); }
    else {
        if (analysis_write(env, erec, ((char *) data) + jpos, jframes)) {
            write_block(env, erec);
        }
    }
}

// The finish method body, also invoked from free.
static void
finish(JNIEnv *env, EncoderRecord *erec) {
    if (erec->finish_called) { return; }
    erec->finish_called = true;
    int err = vorbis_analysis_wrote(&erec->vd, 0);
    if (err < 0) { raiseErrorCodeException(env, err, vorbis_strerror(err)); }
    else { write_block(env, erec); }
    ogg_stream_clear(&erec->os);
    vorbis_block_clear(&erec->vb);
    vorbis_dsp_clear(&erec->vd);
    vorbis_info_clear(&erec->vi);
}

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

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