/* RAU-26, MIT, Copyright (c) 2026 Joerg Burbach.
 * C99 port of Formats/module_format_rau.pbi. No file, UI or audio-device I/O.
 */
#include "rau.h"
#include <stdlib.h>
#include <string.h>
#include <limits.h>

#define MAX_SAMPLES 16777216u
#define FRAME 1024u
static const int fuzzy_table[16] = {0,8,13,21,34,55,89,144,233,288,377,610,850,987,1597,2584};
typedef struct { uint8_t *data; size_t size, capacity; unsigned acc, bits; int error; } writer;
typedef struct { const uint8_t *data; size_t size, pos; unsigned acc, bits; int error; } reader;
typedef struct { int32_t w[2], prev[2]; } lms_state;
static uint32_t le32(const uint8_t *p) { return (uint32_t)p[0] | (uint32_t)p[1]<<8 | (uint32_t)p[2]<<16 | (uint32_t)p[3]<<24; }
static void put32(uint8_t *p, uint32_t v) { for (unsigned i=0; i<4; i++) p[i]=(uint8_t)(v>>(i*8)); }
static int64_t floor_div(int64_t v, int64_t d) { return v>=0 ? v/d : -((-v+d-1)/d); }
static int32_t wrap32(int64_t v) { uint32_t u=(uint32_t)v; return u<=INT32_MAX ? (int32_t)u : (int32_t)((int64_t)u-4294967296LL); }
static int16_t clamp16(int64_t v) { return (int16_t)(v < -32768 ? -32768 : v > 32767 ? 32767 : v); }
static int codec_fuzzy(const rau_info *h) { int f=fuzzy_table[h->fuzzy_index]; return h->bits==8 && f ? (f/256 ? f/256 : 1) : f; }
static int reserve(writer *w, size_t n) {
    if (w->error) return 0;
    if (n > SIZE_MAX-w->size) { w->error=RAU_NOMEM; return 0; }
    if (w->size+n > w->capacity) {
        size_t cap=w->capacity ? w->capacity : 4096;
        while (cap<w->size+n) { if (cap>SIZE_MAX/2) { cap=w->size+n; break; } cap*=2; }
        uint8_t *p=(uint8_t *)realloc(w->data,cap);
        if (!p) { w->error=RAU_NOMEM; return 0; } w->data=p; w->capacity=cap;
    }
    return 1;
}
static void write_bits(writer *w, uint32_t value, unsigned count) {
    for (unsigned i=0; i<count && !w->error; i++) {
        w->acc |= ((value>>i)&1u)<<w->bits++;
        if (w->bits==8) { if (!reserve(w,1)) return; w->data[w->size++]=(uint8_t)w->acc; w->acc=0; w->bits=0; }
    }
}
static uint32_t read_bits(reader *r, unsigned count) {
    uint32_t value=0;
    for (unsigned i=0; i<count && !r->error; i++) {
        if (!r->bits) { if (r->pos>=r->size) { r->error=RAU_TRUNCATED; break; } r->acc=r->data[r->pos++]; r->bits=8; }
        value |= (r->acc&1u)<<i; r->acc>>=1; r->bits--;
    }
    return value;
}
static int32_t read_signed(reader *r, unsigned width) {
    uint32_t v=read_bits(r,width); return (int32_t)(v & (1u<<(width-1)) ? (int64_t)v-(1LL<<width) : v);
}
static uint32_t zigzag(int32_t v) { return v<0 ? (uint32_t)(-(int64_t)v*2-1) : (uint32_t)v*2; }
static uint64_t rice_cost(const int32_t *res, unsigned n, unsigned block, writer *w) {
    uint64_t total=0;
    for (unsigned i=0; i<n; i+=block) {
        unsigned end=i+block<n ? i+block : n, best_k=0; uint64_t best=UINT64_MAX;
        for (unsigned k=0; k<16; k++) {
            uint64_t cost=0;
            for (unsigned j=i; j<end; j++) { uint32_t q=zigzag(res[j])>>k; cost+=q>=48 ? 81 : q+1+k; if (cost>=best) break; }
            if (cost<best) { best=cost; best_k=k; }
        }
        total+=4+best;
        if (w) {
            write_bits(w,best_k,4);
            for (unsigned j=i; j<end; j++) {
                uint32_t z=zigzag(res[j]), q=z>>best_k;
                for (unsigned t=0; t<(q<48 ? q : 48); t++) write_bits(w,1,1);
                write_bits(w,0,1); write_bits(w,q>=48 ? z : z&((1u<<best_k)-1),q>=48 ? 32 : best_k);
            }
        }
    }
    return total;
}
static void read_rice(reader *r, int32_t *res, unsigned n, unsigned block) {
    for (unsigned i=0; i<n && !r->error; i+=block) {
        unsigned k=read_bits(r,4), end=i+block<n ? i+block : n;
        for (unsigned j=i; j<end && !r->error; j++) {
            unsigned q=0;
            while (read_bits(r,1) && !r->error) { if (++q>48) { r->error=RAU_INVALID; break; } }
            uint32_t z=q==48 ? read_bits(r,32) : (q<<k)|read_bits(r,k);
            res[j]=(z&1) ? (int32_t)(-((int64_t)z+1)/2) : (int32_t)(z/2);
        }
    }
}
static int64_t predict(const int32_t *x, unsigned i, unsigned ch, unsigned order) {
    if (order==1) return x[i-ch];
    if (order==2) return 2LL*x[i-ch]-x[i-2*ch];
    if (order==3) return 3LL*x[i-ch]-3LL*x[i-2*ch]+x[i-3*ch];
    return 4LL*x[i-ch]-6LL*x[i-2*ch]+4LL*x[i-3*ch]-x[i-4*ch];
}
static int transform(const int32_t *frame, int32_t *x, unsigned n, unsigned ch, unsigned stereo) {
    int alpha=0;
    if (stereo==2) {
        int64_t ll=0,lr=0;
        for (unsigned i=0; i<n; i+=2) { ll+=(int64_t)frame[i]*frame[i]; lr+=(int64_t)frame[i]*frame[i+1]; }
        if (ll) { int64_t num=8*lr; alpha=(int)(num>=0 ? (num+ll/2)/ll : -((-num+ll/2)/ll)); }
        if (alpha>15) alpha=15; if (alpha< -16) alpha=-16;
    }
    for (unsigned i=0; i<n; i+=ch) {
        if (!stereo || ch==1) { for (unsigned c=0; c<ch; c++) x[i+c]=frame[i+c]; }
        else if (stereo==1) { x[i]=(int32_t)floor_div((int64_t)frame[i]+frame[i+1],2); x[i+1]=frame[i]-frame[i+1]; }
        else { x[i]=frame[i]; x[i+1]=frame[i+1]-(int32_t)floor_div((int64_t)alpha*frame[i],8); }
    }
    return alpha;
}
static void inverse(int32_t *x, unsigned n, unsigned ch, unsigned stereo, int alpha) {
    if (ch==2 && stereo) for (unsigned i=0; i<n; i+=2) {
        if (stereo==2) x[i+1]=wrap32((int64_t)x[i+1]+floor_div((int64_t)alpha*x[i],8));
        else { int32_t r=wrap32((int64_t)x[i]-floor_div(x[i+1],2)); x[i]=wrap32((int64_t)r+x[i+1]); x[i+1]=r; }
    }
}
static uint64_t residual(const int32_t *x, int32_t *res, int32_t *recon, unsigned n, unsigned ch, unsigned order, int fuzzy) {
    uint64_t sad=0; int step=2*fuzzy+1; memcpy(recon,x,order*ch*sizeof(*x));
    for (unsigned i=order*ch; i<n; i++) {
        int64_t p=predict(recon,i,ch,order), r=x[i]-p, q=r;
        if (fuzzy) q=r>=0 ? (r+step/2)/step : -((-r+step/2)/step);
        res[i-order*ch]=wrap32(q); recon[i]=wrap32(p+q*step); sad+=(uint64_t)(q<0 ? -q : q);
    }
    return sad;
}
static void lms(int32_t *res, unsigned n, unsigned ch, lms_state *s, int decode) {
    for (unsigned i=0; i<n; i++) {
        unsigned c=i%ch; int32_t prev=s->prev[c]; int64_t p=floor_div((int64_t)s->w[c]*prev,256);
        int32_t e=decode ? res[i] : wrap32((int64_t)res[i]-p), value=decode ? wrap32(e+p) : res[i];
        res[i]=decode ? value : e;
        s->w[c]+=((e>0)-(e<0))*((prev>0)-(prev<0));
        if (s->w[c]>2047) s->w[c]=2047; if (s->w[c]< -2048) s->w[c]=-2048; s->prev[c]=value;
    }
}
static int parse_header(const uint8_t *data, size_t size, rau_info *h) {
    if (size<17) return RAU_TRUNCATED;
    if (memcmp(data,"RA26",4) || (data[4]&192)) return RAU_INVALID;
    h->bits=data[4]&1 ? 16 : 8; h->channels=((data[4]>>1)&1)+1; h->fuzzy_index=data[4]>>2;
    h->sample_rate=le32(data+5); h->samples=le32(data+9);
    return le32(data+13) || !h->samples || h->samples>MAX_SAMPLES || !h->sample_rate || h->sample_rate>200000 ? RAU_INVALID : RAU_OK;
}
int rau_inspect(const uint8_t *data, size_t size, rau_info *info) {
    if (!info) return RAU_INVALID; memset(info,0,sizeof(*info));
    if (!data || size<4) return RAU_TRUNCATED;
    if (memcmp(data,"RC26",4)) { rau_info h; int err=parse_header(data,size,&h); if (!err) *info=h; return err; }
    size_t pos=4; unsigned packets=0; rau_info first={0}; uint32_t total=0;
    while (pos<size) {
        if (size-pos<4) return RAU_TRUNCATED;
        uint32_t len=le32(data+pos); pos+=4;
        if (len<17 || len>size-pos) return RAU_TRUNCATED;
        rau_info h; int err=parse_header(data+pos,len,&h); if (err) return err;
        if (!packets) first=h;
        else if (h.channels!=first.channels || h.bits!=first.bits || h.fuzzy_index!=first.fuzzy_index || h.sample_rate!=first.sample_rate) return RAU_INVALID;
        if (h.samples>MAX_SAMPLES-total) return RAU_INVALID;
        total+=h.samples; pos+=len; packets++;
    }
    if (!packets) return RAU_INVALID; *info=first; info->samples=total; return RAU_OK;
}
static int decode_packet(const uint8_t *data, size_t size, const rau_info *h, int16_t *pcm) {
    reader br={data+17,size-17,0,0,0,0}; lms_state filter={{0,0},{0,0}};
    int32_t x[FRAME*2],prev[FRAME*2],res[FRAME*2]; unsigned ch=h->channels; int fuzzy=codec_fuzzy(h);
    for (unsigned start=0; start<h->samples; start+=FRAME) {
        unsigned count=h->samples-start<FRAME ? h->samples-start : FRAME, n=count*ch;
        if (read_bits(&br,2)!=3) return br.error ? br.error : RAU_INVALID;
        unsigned kind=read_bits(&br,5);
        if (kind==1) { int32_t v[2]; for (unsigned c=0; c<ch; c++) v[c]=read_signed(&br,h->bits); for (unsigned i=0; i<n; i++) x[i]=v[i%ch]; }
        else if (kind==2) { if (!start) return RAU_INVALID; memcpy(x,prev,n*sizeof(*x)); }
        else if (kind==0) {
            unsigned order=read_bits(&br,2)+1, skip=read_bits(&br,1), stereo=ch==2 ? read_bits(&br,2) : 0;
            int alpha=stereo==2 ? read_signed(&br,5) : 0;
            if (count<=order || stereo>2) return RAU_INVALID;
            for (unsigned i=0; i<order*ch; i++) x[i]=read_signed(&br,h->bits+(ch==2 && i%ch==1));
            unsigned values=(count-order)*ch, block=64u<<read_bits(&br,2); read_rice(&br,res,values,block);
            if (br.error) return br.error;
            if (!fuzzy && !skip) lms(res,values,ch,&filter,1);
            for (unsigned i=order*ch; i<n; i++) x[i]=wrap32(predict(x,i,ch,order)+(int64_t)res[i-order*ch]*(2*fuzzy+1));
            inverse(x,n,ch,stereo,alpha);
        } else return br.error ? br.error : RAU_INVALID;
        if (br.error) return br.error;
        for (unsigned i=0; i<n; i++) pcm[start*ch+i]=clamp16((int64_t)x[i]*(h->bits==8 ? 256 : 1));
        memcpy(prev,x,n*sizeof(*x));
    }
    return RAU_OK;
}
int rau_decode(const uint8_t *data, size_t size, rau_audio *audio) {
    if (!audio) return RAU_INVALID; memset(audio,0,sizeof(*audio)); rau_info info;
    int err=rau_inspect(data,size,&info); if (err) return err;
    int16_t *pcm=(int16_t *)malloc((size_t)info.samples*info.channels*sizeof(*pcm)); if (!pcm) return RAU_NOMEM;
    if (!memcmp(data,"RC26",4)) {
        size_t pos=4, offset=0;
        while (pos<size) { uint32_t len=le32(data+pos); pos+=4; rau_info h; parse_header(data+pos,len,&h);
            err=decode_packet(data+pos,len,&h,pcm+offset); if (err) break; offset+=(size_t)h.samples*h.channels; pos+=len; }
    } else err=decode_packet(data,size,&info,pcm);
    if (err) { free(pcm); return err; }
    audio->pcm=pcm; audio->sample_rate=info.sample_rate; audio->samples=info.samples; audio->channels=info.channels; audio->bits=16; return RAU_OK;
}
static int encode_packet(const int16_t *pcm, const rau_info *h, writer *w) {
    if (!reserve(w,17)) return w->error;
    memcpy(w->data,"RA26",4); w->data[4]=(uint8_t)((h->bits==16) | ((h->channels-1)<<1) | (h->fuzzy_index<<2));
    put32(w->data+5,h->sample_rate); put32(w->data+9,h->samples); put32(w->data+13,0); w->size=17;
    int32_t frame[FRAME*2],prev[FRAME*2],x[FRAME*2],res[FRAME*2],recon[FRAME*2],trial[FRAME*2];
    lms_state filter={{0,0},{0,0}}; unsigned ch=h->channels; int fuzzy=codec_fuzzy(h);
    for (unsigned start=0; start<h->samples && !w->error; start+=FRAME) {
        unsigned count=h->samples-start<FRAME ? h->samples-start : FRAME, n=count*ch; int constant=1,repeat=start>0;
        for (unsigned i=0; i<n; i++) frame[i]=h->bits==8 ? (int32_t)floor_div(pcm[start*ch+i],256) : pcm[start*ch+i];
        for (unsigned i=0; i<n; i++) { if (abs(frame[i]-frame[i%ch])>fuzzy) constant=0; if (start && abs(frame[i]-prev[i])>fuzzy) repeat=0; }
        unsigned stereo_best=0,order_best=0,mode=0; uint64_t sad=UINT64_MAX,cost=UINT64_MAX;
        for (unsigned stereo=0; stereo<=(ch==2 ? 2u : 0u); stereo++) {
            transform(frame,x,n,ch,stereo);
            for (unsigned order=1; order<=(fuzzy ? 2u : 3u) && order<count; order++) {
                uint64_t score=residual(x,res,recon,n,ch,order,fuzzy);
                if (score<sad) { sad=score; stereo_best=stereo; order_best=order; }
            }
        }
        int alpha=0;
        if (order_best) {
            alpha=transform(frame,x,n,ch,stereo_best); residual(x,res,recon,n,ch,order_best,fuzzy);
            for (unsigned m=0; m<4; m++) { uint64_t c=rice_cost(res,n-order_best*ch,64u<<m,NULL); if (c<cost) { cost=c; mode=m; } }
            cost+=12+(ch==2 ? 2+(stereo_best==2 ? 5 : 0) : 0)+order_best*(ch*h->bits+(ch==2));
        }
        uint64_t constant_cost=7+ch*h->bits;
        write_bits(w,3,2);
        if (repeat && 7<(constant && constant_cost<cost ? constant_cost : cost)) write_bits(w,2,5);
        else if (constant && constant_cost<=cost) {
            write_bits(w,1,5); for (unsigned c=0; c<ch; c++) write_bits(w,(uint32_t)frame[c],h->bits);
            for (unsigned i=0; i<n; i++) prev[i]=frame[i%ch];
        } else {
            if (!order_best) return RAU_INVALID;
            unsigned values=n-order_best*ch,skip=1; write_bits(w,0,5); write_bits(w,order_best-1,2);
            if (!fuzzy) { lms_state saved=filter; memcpy(trial,res,values*sizeof(*res)); lms(trial,values,ch,&filter,0);
                if (rice_cost(trial,values,64u<<mode,NULL)<rice_cost(res,values,64u<<mode,NULL)) { memcpy(res,trial,values*sizeof(*res)); skip=0; } else filter=saved; }
            write_bits(w,skip,1); if (ch==2) { write_bits(w,stereo_best,2); if (stereo_best==2) write_bits(w,(uint32_t)alpha,5); }
            for (unsigned i=0; i<order_best*ch; i++) write_bits(w,(uint32_t)x[i],h->bits+(ch==2 && i%ch==1));
            write_bits(w,mode,2); rice_cost(res,values,64u<<mode,w); inverse(recon,n,ch,stereo_best,alpha); memcpy(prev,recon,n*sizeof(*prev));
        }
    }
    if (w->bits) write_bits(w,0,8-w->bits); return w->error;
}
static int16_t sample(const rau_audio *a, int64_t i, unsigned c) {
    if (i<0) i=0; if (i>=a->samples) i=a->samples-1; return a->pcm[(size_t)i*a->channels+c];
}
int rau_encode(const rau_audio *audio, rau_profile profile, uint8_t **data, size_t *size) {
    if (data) *data=NULL; if (size) *size=0;
    if (!data || !size || !audio || !audio->pcm || !audio->samples || audio->samples>MAX_SAMPLES ||
        audio->channels<1 || audio->channels>2 || !audio->sample_rate || audio->sample_rate>200000 ||
        (audio->bits!=8 && audio->bits!=16) || profile<RAU_LOSSLESS || profile>RAU_CRAPPY) return RAU_INVALID;
    static const uint32_t rates[5]={0,16000,24000,32000,8000};
    static const unsigned codes[5][2]={{0,0},{10,11},{7,8},{4,6},{12,14}};
    uint32_t rate=profile==RAU_LOSSLESS || audio->sample_rate<rates[profile] ? audio->sample_rate : rates[profile];
    uint32_t count=(uint32_t)(((uint64_t)audio->samples*rate+audio->sample_rate-1)/audio->sample_rate);
    unsigned ch=audio->channels; int16_t *prepared=NULL; const int16_t *pcm=audio->pcm;
    if (rate!=audio->sample_rate) {
        prepared=(int16_t *)malloc((size_t)count*ch*sizeof(*prepared)); if (!prepared) return RAU_NOMEM;
        for (uint32_t i=0; i<count; i++) {
            uint64_t pos=(uint64_t)i*audio->sample_rate*65536/rate; int64_t base=(int64_t)(pos>>16); unsigned frac=(unsigned)(pos&65535);
            if (base>=audio->samples-1) { base=audio->samples-1; frac=0; }
            for (unsigned c=0; c<ch; c++) {
                int64_t a=floor_div(sample(audio,base-1,c)+2LL*sample(audio,base,c)+sample(audio,base+1,c),4);
                int64_t b0=base+1<audio->samples ? sample(audio,base+1,c) : a;
                int64_t b=floor_div(sample(audio,base,c)+2*b0+(base+2<audio->samples ? sample(audio,base+2,c) : b0),4);
                prepared[(size_t)i*ch+c]=(int16_t)(a+floor_div((b-a)*frac,65536));
            }
        }
        pcm=prepared;
    }
    rau_info h={rate,0,ch,profile==RAU_CRAPPY ? 8u : profile==RAU_LOSSLESS ? audio->bits : 16u,codes[profile][ch-1]};
    writer out={0}; int err=RAU_OK;
    if (!reserve(&out,4)) { free(prepared); return out.error; } memcpy(out.data,"RC26",4); out.size=4;
    for (uint32_t start=0; start<count; start+=rate) {
        h.samples=count-start<rate ? count-start : rate; writer packet={0}; err=encode_packet(pcm+(size_t)start*ch,&h,&packet);
        if (!err && reserve(&out,4+packet.size)) { put32(out.data+out.size,(uint32_t)packet.size); out.size+=4; memcpy(out.data+out.size,packet.data,packet.size); out.size+=packet.size; }
        if (!err) err=out.error; free(packet.data); if (err) break;
    }
    free(prepared); if (err) { free(out.data); return err; } *data=out.data; *size=out.size; return RAU_OK;
}
