/* SPDX-License-Identifier: MIT. PMF0/PMFZ-26 codec. Link RAU and shared BriefLZ/LeXA. */
#include "pmf0.h"
#include "../../rau/codecs/rau.h"
#include "../../shared/brieflz-container.h"
#include <stdlib.h>
#include <string.h>
#define LIMIT ((size_t)67108864)
typedef struct {const uint8_t *s;size_t n,p;int error;} reader;
typedef struct {uint8_t *s;size_t n,capacity;int error;} writer;
typedef struct {const uint8_t *s;size_t n,bit;unsigned k;int error;} bitreader;
static unsigned u16(const uint8_t *s){return s[0]|(unsigned)s[1]<<8;}
static uint32_t u32(const uint8_t *s){return s[0]|(uint32_t)s[1]<<8|(uint32_t)s[2]<<16|(uint32_t)s[3]<<24;}
static void put(uint8_t *p,size_t v,unsigned n){unsigned i;for(i=0;i<n;i++)p[i]=(uint8_t)(v>>(i*8));}
static const uint8_t *take(reader *r,size_t n){const uint8_t *p;if(r->error||r->p>r->n||n>r->n-r->p){r->error=1;return NULL;}p=r->s+r->p;r->p+=n;return p;}
static unsigned get(reader *r,unsigned n){const uint8_t *p=take(r,n);return !p?0:n==1?p[0]:n==2?u16(p):u32(p);}
static void write_bytes(writer *w,const uint8_t *s,size_t n){uint8_t *p;size_t capacity;if(w->error||n>LIMIT-w->n){w->error=1;return;}if(n>w->capacity-w->n){capacity=w->capacity?w->capacity:1024;while(capacity<w->n+n)capacity=capacity>LIMIT/2?LIMIT:capacity*2;p=realloc(w->s,capacity);if(!p){w->error=1;return;}w->s=p;w->capacity=capacity;}if(n)memcpy(w->s+w->n,s,n);w->n+=n;}
static void write_uint(writer *w,size_t v,unsigned n){uint8_t s[4];put(s,v,n);write_bytes(w,s,n);}
static int raw_info(const uint8_t *s,size_t n,pmf0_info *o,size_t *audio_size){size_t i,expected,total=0,off,len,a;unsigned flags;
    if(!s||n<48||n>LIMIT||memcmp(s,"PMF0",4)||s[4]!=26||(s[5]!=11&&s[5]!=12))return -1;
    memset(o,0,sizeof(*o));o->profile=s[5];o->root=(uint16_t)u16(s+6);o->min=(uint16_t)u16(s+8);o->width=u32(s+10);o->height=u32(s+14);o->count=u32(s+18);o->fps=(uint16_t)u16(s+22);o->gop=(uint16_t)u16(s+24);flags=s[30];o->audio_rate=u32(s+40);o->audio_channels=s[44];o->audio_bits=s[45];
    if(o->root<8||o->root>1024||(o->root&(o->root-1))||o->min!=8||!o->width||!o->height||o->width>8192||o->height>8192||(size_t)o->width*o->height>16777216||!o->count||!o->fps||!o->gop||u32(s+26)!=48||flags>1||s[31]||u32(s+36)||u16(s+46)||o->count>(n-48)/12)return -1;
    if(flags){if(!o->audio_rate||(o->audio_channels!=1&&o->audio_channels!=2)||(o->audio_bits!=8&&o->audio_bits!=16))return -1;}else if(o->audio_rate||o->audio_channels||o->audio_bits)return -1;
    expected=48+(size_t)o->count*12;for(i=0;i<o->count;i++){off=u32(s+48+i*12);len=u32(s+52+i*12);a=u32(s+56+i*12);if(off!=expected||off>n||!len||len>n-off||a>n-off-len)return -1;expected+=len+a;total+=a;}
    if(expected!=n||total!=u32(s+32)||!!total!=!!flags)return -1;if(audio_size)*audio_size=total;return 0;
}
static unsigned bit(bitreader *b){unsigned v;if(b->error||b->bit>=b->n*8){b->error=1;return 0;}v=b->s[b->bit/8]>>(b->bit%8)&1;b->bit++;return v;}
static int rice(bitreader *b){unsigned q=0,v=0,i,z;while(!bit(b)&&!b->error)if(++q>1048576){b->error=1;break;}for(i=0;i<b->k;i++)v|=bit(b)<<i;z=(q<<b->k)|v;return z&1?-(int)((z+1)/2):(int)(z/2);}
static int paeth(int l,int u,int c){int p=l+u-c,a=abs(p-l),b=abs(p-u),d=abs(p-c);return a<=b&&a<=d?l:b<=d?u:c;}
static int sample(const uint8_t *a,int x,int y,unsigned c,unsigned W,unsigned H){return a&&x>=0&&y>=0&&(unsigned)x<W&&(unsigned)y<H?a[((size_t)y*W+(unsigned)x)*3+c]:0;}
static int node(reader *r,uint8_t *cur,uint8_t **history,unsigned frame,const pmf0_info *o,unsigned bx,unsigned by,unsigned size,unsigned base_x,unsigned base_y){
    unsigned tag=get(r,1),half,bw,bh,x,y,c,lag=0,sx,sy,mode,mask,ks[3]={0,0,0};int dx=0,dy=0,value,pred,l,u,corner;size_t lens[3]={0,0,0};const uint8_t *ref=NULL,*colours,*map;bitreader planes[3];
    if(r->error)return -1;if(!tag){if(size<=o->min)return -1;half=size/2;return node(r,cur,history,frame,o,bx,by,half,base_x,base_y)||node(r,cur,history,frame,o,bx+half,by,half,base_x,base_y)||node(r,cur,history,frame,o,bx,by+half,half,base_x,base_y)||node(r,cur,history,frame,o,bx+half,by+half,half,base_x,base_y)?-1:0;}
    if(tag>29)return -1;if(bx>=o->width||by>=o->height){if(tag!=1)return -1;get(r,2);get(r,2);return r->error?-1:0;}bw=o->width-bx<size?o->width-bx:size;bh=o->height-by<size?o->height-by:size;
    if(tag==1||tag==2){if(tag==2)lag=get(r,1);sx=get(r,2);sy=get(r,2);ref=tag==2?history[(frame-lag)&7]:cur;if((tag==2&&(!lag||lag>8||lag>frame))||!ref||sx>o->width||bw>o->width-sx||sy>o->height||bh>o->height-sy||(tag==1&&(sy>by||(sy==by&&sx>=bx)))||r->error)return -1;for(y=0;y<bh;y++)for(x=0;x<bw;x++)for(c=0;c<3;c++)cur[((by+y)*o->width+bx+x)*3+c]=ref[((sy+y)*o->width+sx+x)*3+c];return 0;}
    if(tag==3||tag==28||tag==29){unsigned bits=tag==28?1:2;size_t row=(size*bits+7)/8;colours=take(r,tag==3?3:tag==28?6:12);map=tag==3?NULL:take(r,size*row);if(r->error)return -1;for(y=0;y<bh;y++)for(x=0;x<bw;x++){unsigned index=map?map[y*row+x*bits/8]>>(x*bits%8)&((1u<<bits)-1):0;memcpy(cur+((by+y)*o->width+bx+x)*3,colours+index*3,3);}return 0;}
    mode=tag>=20?2:tag>=12?1:0;mask=tag-(mode==2?20:mode==1?12:4);if(mask>7)return -1;if(mode)lag=get(r,1);if(mode==2){dx=(int16_t)get(r,2);dy=(int16_t)get(r,2);}if(mode){if(!lag||lag>8||lag>frame)return -1;ref=history[(frame-lag)&7];if(!ref)return -1;}
    memset(planes,0,sizeof(planes));for(c=0;c<3;c++)if(mask&(1u<<c)){ks[c]=get(r,1);lens[c]=get(r,1);if(lens[c]==255)lens[c]=get(r,4);if(ks[c]>15)return -1;}
    for(c=0;c<3;c++){planes[c].s=lens[c]?take(r,lens[c]):NULL;planes[c].n=lens[c];planes[c].k=ks[c];}if(r->error)return -1;
    for(y=by;y<by+bh;y++)for(x=bx;x<bx+bw;x++)for(c=0;c<3;c++){
        l=x>base_x?sample(cur,(int)x-1,(int)y,c,o->width,o->height)-sample(ref,(int)x-1+dx,(int)y+dy,c,o->width,o->height):0;
        u=y>base_y?sample(cur,(int)x,(int)y-1,c,o->width,o->height)-sample(ref,(int)x+dx,(int)y-1+dy,c,o->width,o->height):0;
        corner=x>base_x&&y>base_y?sample(cur,(int)x-1,(int)y-1,c,o->width,o->height)-sample(ref,(int)x-1+dx,(int)y-1+dy,c,o->width,o->height):0;
        pred=x<=base_x&&y<=base_y?0:x<=base_x?u:y<=base_y?l:paeth(l,u,corner);value=pred+sample(ref,(int)x+dx,(int)y+dy,c,o->width,o->height)+(lens[c]?rice(&planes[c]):0);if(value<0||value>255||planes[c].error)return -1;cur[((size_t)y*o->width+x)*3+c]=(uint8_t)value;
    }
    return 0;
}
void pmf0_free(pmf0_video *v){if(v){free(v->rgba);free(v->audio);memset(v,0,sizeof(*v));}}
static int decode_raw(const uint8_t *s,size_t n,pmf0_video *out){
    uint8_t *history[8]={NULL,NULL,NULL,NULL,NULL,NULL,NULL,NULL},*cur=NULL;size_t audio_size,frame_size,pixels,i,j,root,roots,columns,off,len,start,end,audio_pos=0;reader r;int rc=-1;
    memset(out,0,sizeof(*out));if(raw_info(s,n,&out->info,&audio_size))return -1;pixels=(size_t)out->info.width*out->info.height;frame_size=pixels*4;if(out->info.count>LIMIT/frame_size)return -1;
    out->rgba=malloc(frame_size*out->info.count);out->audio=malloc(audio_size?audio_size:1);if(!out->rgba||!out->audio)goto done;out->audio_size=audio_size;columns=(out->info.width+out->info.root-1)/out->info.root;roots=columns*((out->info.height+out->info.root-1)/out->info.root);
    for(i=0;i<out->info.count;i++){
        off=u32(s+48+i*12);len=u32(s+52+i*12);if(roots>len/4)goto done;cur=calloc(pixels*3,1);if(!cur)goto done;start=off+roots*4;
        for(root=0;root<roots;root++){if(u32(s+off+root*4)!=start)goto done;end=root+1<roots?u32(s+off+(root+1)*4):off+len;if(end<=start||end>off+len)goto done;r.s=s+start;r.n=end-start;r.p=0;r.error=0;if(node(&r,cur,history,(unsigned)i,&out->info,(unsigned)(root%columns)*out->info.root,(unsigned)(root/columns)*out->info.root,out->info.root,(unsigned)(root%columns)*out->info.root,(unsigned)(root/columns)*out->info.root)||r.error||r.p!=r.n)goto done;start=end;}
        free(history[i&7]);history[i&7]=cur;cur=NULL;
        for(j=0;j<pixels;j++){int a=history[i&7][j*3],b=history[i&7][j*3+1],c=history[i&7][j*3+2],values[3];unsigned channel;values[0]=out->info.profile==12?a+b-c:a;values[1]=out->info.profile==12?a+c-128:b;values[2]=out->info.profile==12?a-b-c+256:c;for(channel=0;channel<3;channel++)out->rgba[i*frame_size+j*4+channel]=(uint8_t)(values[channel]<0?0:values[channel]>255?255:values[channel]);out->rgba[i*frame_size+j*4+3]=255;}
        len=u32(s+56+i*12);if(len){memcpy(out->audio+audio_pos,s+start,len);audio_pos+=len;}
    }
    if(audio_size){rau_info ai;if(rau_inspect(out->audio,audio_size,&ai)||ai.sample_rate!=out->info.audio_rate||ai.channels!=out->info.audio_channels||ai.bits!=out->info.audio_bits)goto done;}rc=0;
done:free(cur);for(i=0;i<8;i++)free(history[i]);if(rc)pmf0_free(out);return rc;
}
int pmf0_decode(const uint8_t *s,size_t n,pmf0_video *out){
    size_t count,total,expected=0,pos,i,raw_size,stored;unsigned frames,method;uint8_t *raw=NULL;pmf0_video part;int rc=-1;
    if(!out)return -1;memset(out,0,sizeof(*out));if(!s||n>LIMIT)return -1;if(n>=4&&!memcmp(s,"PMF0",4))return decode_raw(s,n,out);
    if(n<20||memcmp(s,"PMFZ",4)||s[4]!=26||s[5]||u16(s+6)!=25||!u32(s+8)||!u16(s+12)||u32(s+16)!=20)return -1;total=u32(s+8);count=u16(s+14);if(count!=(total+24)/25||count>(n-20)/16)return -1;pos=20+count*16;
    for(i=0;i<count;i++){
        const uint8_t *entry=s+20+i*16;frames=u16(entry+4);method=entry[6];raw_size=u32(entry+8);stored=u32(entry+12);
        if(u32(entry)!=expected||!frames||frames>25||entry[7]||(method!=0&&method!=14)||raw_size<48||raw_size>LIMIT||!stored||pos>n||stored>n-pos)goto done;
        if(method)raw=codec_brief_lz_decode(s+pos,stored,raw_size);else{if(stored!=raw_size)goto done;raw=malloc(raw_size);if(raw)memcpy(raw,s+pos,raw_size);}if(!raw)goto done;
        if(decode_raw(raw,raw_size,&part)){free(raw);raw=NULL;goto done;}free(raw);raw=NULL;
        if(part.info.count!=frames||part.info.fps!=u16(s+12)||part.audio_size){pmf0_free(&part);goto done;}
        if(!i){size_t frame_size=(size_t)part.info.width*part.info.height*4;if(total>LIMIT/frame_size){pmf0_free(&part);goto done;}out->info=part.info;out->info.count=(uint32_t)total;out->rgba=malloc(total*frame_size);if(!out->rgba){pmf0_free(&part);goto done;}}
        else if(part.info.width!=out->info.width||part.info.height!=out->info.height||part.info.profile!=out->info.profile){pmf0_free(&part);goto done;}
        if(frames>total-expected){pmf0_free(&part);goto done;}memcpy(out->rgba+expected*(size_t)out->info.width*out->info.height*4,part.rgba,(size_t)frames*out->info.width*out->info.height*4);pmf0_free(&part);expected+=frames;pos+=stored;
    }if(expected!=total||pos!=n)goto done;rc=0;
done:free(raw);if(rc)pmf0_free(out);return rc;
}
int pmf0_inspect(const uint8_t *s,size_t n,pmf0_info *o){pmf0_video v;size_t audio;if(!o)return -1;if(n>=4&&s&&!memcmp(s,"PMF0",4))return raw_info(s,n,o,&audio);if(pmf0_decode(s,n,&v))return -1;*o=v.info;pmf0_free(&v);return 0;}
static void write_bit(writer *w,unsigned *bit_pos,unsigned value){if(!*bit_pos)write_uint(w,0,1);if(!w->error&&value)w->s[w->n-1]|=(uint8_t)(1u<<*bit_pos);*bit_pos=(*bit_pos+1)&7;}
int pmf0_encode(const pmf0_video *v,unsigned quality,uint8_t **data,size_t *size){
    writer out={NULL,0,0,0},tree={NULL,0,0,0},planes[3];size_t pixels,frame_size,columns,roots,frame,root,table_pos,frame_start,roots_pos,i;unsigned W,H,x,y,c,bx,by,bit_pos,shift;uint8_t *rgb=NULL;rau_info ai;
    if(!v||!data||!size||!v->rgba||quality>4)return -1;*data=NULL;*size=0;W=v->info.width;H=v->info.height;if(!W||!H||W>8192||H>8192||(size_t)W*H>16777216||!v->info.count||!v->info.fps)return -1;pixels=(size_t)W*H;frame_size=pixels*4;if(v->info.count>LIMIT/frame_size||v->audio_size>LIMIT)return -1;
    memset(&ai,0,sizeof(ai));if(v->audio_size&&rau_inspect(v->audio,v->audio_size,&ai))return -1;shift=quality==4?0:quality==3?2:quality==2?3:quality==1?4:5;columns=(W+7)/8;roots=columns*((H+7)/8);rgb=malloc(pixels*3);if(!rgb)return -1;
    write_bytes(&out,(const uint8_t *)"PMF0",4);write_uint(&out,26,1);write_uint(&out,11,1);write_uint(&out,8,2);write_uint(&out,8,2);write_uint(&out,W,4);write_uint(&out,H,4);write_uint(&out,v->info.count,4);write_uint(&out,v->info.fps,2);write_uint(&out,1,2);write_uint(&out,48,4);write_uint(&out,!!v->audio_size,1);write_uint(&out,0,1);write_uint(&out,v->audio_size,4);write_uint(&out,0,4);write_uint(&out,ai.sample_rate,4);write_uint(&out,ai.channels,1);write_uint(&out,ai.bits,1);write_uint(&out,0,2);table_pos=out.n;for(i=0;i<(size_t)v->info.count*3;i++)write_uint(&out,0,4);
    for(frame=0;frame<v->info.count&&!out.error;frame++){
        for(i=0;i<pixels*3;i++)rgb[i]=(uint8_t)(v->rgba[frame*frame_size+i/3*4+i%3]>>shift<<shift);frame_start=out.n;roots_pos=out.n;for(i=0;i<roots;i++)write_uint(&out,0,4);
        for(root=0;root<roots&&!out.error;root++){
            bx=(unsigned)(root%columns)*8;by=(unsigned)(root/columns)*8;memset(planes,0,sizeof(planes));tree.n=0;tree.error=0;
            for(c=0;c<3;c++){bit_pos=0;for(y=by;y<H&&y<by+8;y++)for(x=bx;x<W&&x<bx+8;x++){size_t off=((size_t)y*W+x)*3+c;int l=x>bx?rgb[off-3]:0,u=y>by?rgb[off-W*3]:0,ul=x>bx&&y>by?rgb[off-W*3-3]:0,p=x==bx&&y==by?0:x==bx?u:y==by?l:paeth(l,u,ul),d=rgb[off]-p;unsigned z=d<0?(unsigned)(-d*2-1):(unsigned)d*2,b;for(b=0;b<(z>>8);b++)write_bit(&planes[c],&bit_pos,0);write_bit(&planes[c],&bit_pos,1);for(b=0;b<8;b++)write_bit(&planes[c],&bit_pos,z>>b&1);}}
            write_uint(&tree,11,1);for(c=0;c<3;c++){write_uint(&tree,8,1);write_uint(&tree,planes[c].n<255?planes[c].n:255,1);if(planes[c].n>=255)write_uint(&tree,planes[c].n,4);}for(c=0;c<3;c++){if(planes[c].error)tree.error=1;write_bytes(&tree,planes[c].s,planes[c].n);free(planes[c].s);}
            if(tree.error){out.error=1;break;}if(!out.error)put(out.s+roots_pos+root*4,out.n,4);write_bytes(&out,tree.s,tree.n);
        }
        if(!out.error){put(out.s+table_pos+frame*12,frame_start,4);put(out.s+table_pos+frame*12+4,out.n-frame_start,4);put(out.s+table_pos+frame*12+8,frame?0:v->audio_size,4);}if(!frame)write_bytes(&out,v->audio,v->audio_size);
    }
    free(rgb);free(tree.s);if(out.error){free(out.s);return -1;}*data=out.s;*size=out.n;return 0;
}
