#!/usr/bin/env python3
"""CS10 teaching codecs. Python standard library only; NOT a DEFLATE implementation.
HUF1/ARC1 are local teaching formats, MSB-first bit packing, big-endian fields.
Run without arguments for generated endpoint vectors and exhaustive/negative tests.
Assertions below use explicit exceptions, so python -O runs the same tests.
"""
from collections import Counter
from fractions import Fraction
from itertools import product
import heapq, json, struct
EOF = 256
class Invalid(ValueError): pass

def need(ok, why):
    if not ok: raise Invalid(why)

def equal(a,b,why):
    if a != b: raise RuntimeError(f'{why}: {a!r} != {b!r}')

def rejected(fn):
    try: fn()
    except Invalid: return True
    raise RuntimeError('malformed input accepted')

def pack_bits(bits):
    need(all(b in '01' for b in bits),'not bits')
    pad=(-len(bits))%8
    s=bits+'0'*pad
    return bytes(int(s[i:i+8],2) for i in range(0,len(s),8)),pad

def unpack_bits(blob,count):
    need(type(count) is int and count>=0,'negative bit count')
    need(len(blob)==(count+7)//8,'physical payload length')
    allbits=''.join(f'{v:08b}' for v in blob)
    need('1' not in allbits[count:],'nonzero byte padding')
    return allbits[:count]

def huffman_lengths(data):
    freq=Counter(data);freq[EOF]=1
    heap=[(f,s,{s:0}) for s,f in sorted(freq.items())];heapq.heapify(heap)
    if len(heap)==1:return {EOF:1}
    while len(heap)>1:
        f,a,x=heapq.heappop(heap);g,b,y=heapq.heappop(heap)
        z={s:l+1 for s,l in (x|y).items()}
        heapq.heappush(heap,(f+g,min(a,b),z))
    return heap[0][2]

def canonical(lengths,max_bits=15):
    need(1<=len(lengths)<=257,'table size')
    need(all(type(s) is int and 0<=s<=EOF for s in lengths),'symbol range')
    need(all(type(l) is int and 1<=l<=max_bits for l in lengths.values()),'code length')
    counts=[0]*(max_bits+1)
    for l in lengths.values(): counts[l]+=1
    left=1
    for l in range(1,max_bits+1):
        left=2*left-counts[l]
        need(left>=0,'oversubscribed lengths')
    code=0;next_code=[0]*(max_bits+1)
    for l in range(1,max_bits+1):
        code=(code+counts[l-1])*2;next_code[l]=code
    result={}
    for s,l in sorted(lengths.items()):
        result[s]=f'{next_code[l]:0{l}b}';next_code[l]+=1
    return result

def canonical_reference(lengths):
    # Independent small-instance enumeration: choose the first free binary word.
    chosen={}
    for s,l in sorted(lengths.items(),key=lambda x:(x[1],x[0])):
        for k in range(1<<l):
            word=f'{k:0{l}b}'
            if not any(word.startswith(c) for c in chosen.values()):
                chosen[s]=word;break
        else: raise Invalid('no free leaf')
    return chosen

def huf_encode(data):
    lengths=huffman_lengths(data);codes=canonical(lengths)
    bits=''.join(codes[s] for s in list(data)+[EOF]);payload,pad=pack_bits(bits)
    header=b'HUF1'+struct.pack('>HII',len(lengths),len(data),len(bits))
    header+=b''.join(struct.pack('>HB',s,l) for s,l in sorted(lengths.items()))
    return header+payload, {'lengths':lengths,'codes':codes,'bits':bits,'padding_bits':pad,'header_bytes':len(header)}

def parse_header(blob,magic,max_output,max_bits):
    need(all(type(v) is int and v>=0 for v in (max_output,max_bits)),'nonnegative integer budgets')
    need(len(blob)>=4 and blob[:4]==magic,'magic')
    if magic==b'HUF1':
        need(len(blob)>=14,'short header')
        m,n,bits=struct.unpack('>HII',blob[4:14]);start=14;size=3;w=None
    else:
        need(len(blob)>=15,'short header')
        w=blob[4];need(w in (8,16),'word width')
        m,n,bits=struct.unpack('>HII',blob[5:15]);start=15;size=4
    # Validate attacker-controlled sizes before loops, output allocation, or bit expansion.
    need(1<=m<=257,'table size');need(n<=max_output,'output budget');need(bits<=max_bits,'input bit budget')
    end=start+size*m;need(len(blob)>=end,'short table')
    pairs=[]
    for pos in range(start,end,size):
        s,v=struct.unpack('>HB' if size==3 else '>HH',blob[pos:pos+size]);pairs.append((s,v))
    need(all(0<=s<=EOF for s,_ in pairs),'symbol range')
    need(all(pairs[i][0]<pairs[i+1][0] for i in range(m-1)),'symbol order or duplicate')
    table=dict(pairs);need(EOF in table,'missing EOF')
    return w,n,bits,table,blob[end:]

def huf_decode(blob,max_output=100000,max_bits=1000000):
    _,n,count,lengths,payload=parse_header(blob,b'HUF1',max_output,max_bits)
    need(len(lengths)!=1 or lengths=={EOF:1},'singleton EOF length must be one')
    codes=canonical(lengths);inverse={v:k for k,v in codes.items()};bits=unpack_bits(payload,count)
    out=bytearray();word='';trace=[];ended=False
    for pos,b in enumerate(bits,1):
        word+=b
        if word in inverse:
            s=inverse[word];trace.append({'symbol':s,'word':word,'end_bit':pos});word=''
            if s==EOF:
                need(pos==len(bits),'data after EOF');ended=True;break
            need(len(out)<n and len(out)<max_output,'decoded output exceeds bound');out.append(s)
        else:need(len(word)<max(lengths.values()),'invalid or truncated codeword')
    need(ended and not word and len(out)==n,'missing EOF or output length mismatch')
    return bytes(out),trace

def huf_decode_reference(bits,lengths):
    tree={}
    for s,word in canonical_reference(lengths).items():
        node=tree
        for b in word:node=node.setdefault(b,{})
        node['symbol']=s
    out=[];node=tree
    for b in bits:
        need(b in node,'absent branch');node=node[b]
        if 'symbol' in node:out.append(node['symbol']);node=tree
    need(node is tree and out and out[-1]==EOF,'unfinished reference decode')
    return bytes(out[:-1])

class Model:
    """Stable symbol order. Update AFTER coding/decoding a non-EOF symbol.
    If total reaches threshold: ceil-halving all counts (keeps every count positive).
    """
    def __init__(self,freq,adaptive=False,threshold=None):
        self.freq=dict(sorted(freq.items()));self.adaptive=adaptive;self.threshold=threshold
        need(EOF in self.freq and all(type(s) is int and 0<=s<=EOF for s in self.freq),'model symbols')
        need(all(type(f) is int and f>0 for f in self.freq.values()),'nonpositive frequency')
        if adaptive:need(len(freq)<threshold and sum(freq.values())<threshold,'adaptive threshold')
    def ranges(self):
        lo=0;out=[]
        for s,f in self.freq.items():out.append((s,lo,lo+f));lo+=f
        return out,lo
    def update(self,s):
        if not self.adaptive or s==EOF:return
        self.freq[s]+=1
        if sum(self.freq.values())>=self.threshold:
            self.freq={x:(f+1)//2 for x,f in self.freq.items()}

def arithmetic_encode(data,freq,w=8,adaptive=False,threshold=None):
    need(w in (8,16),'word width');Q=1<<(w-2);H=2*Q;TOP=4*Q-1
    model=Model(freq,adaptive,threshold);need(sum(freq.values())<=Q-1,'frequency total')
    if adaptive:need(threshold<=Q-1,'frequency threshold')
    low=0;high=TOP;pending=0;out=[];trace=[]
    def emit(bit):
        nonlocal pending
        out.append(str(bit));out.extend(str(1-bit) for _ in range(pending));pending=0
    for s in list(data)+[EOF]:
        out_start=len(out)
        ranges,total=model.ranges();matches=[(a,b) for x,a,b in ranges if x==s];need(matches,'unknown symbol')
        a,b=matches[0];R=high-low+1;old=(low,high)
        high=low+R*b//total-1;low=low+R*a//total;narrowed=(low,high);ops=[]
        need(low<=high,'empty integer interval')
        while True:
            if high<H:emit(0);ops.append('E1')
            elif low>=H:emit(1);low-=H;high-=H;ops.append('E2')
            elif low>=Q and high<3*Q:pending+=1;low-=Q;high-=Q;ops.append('E3')
            else:break
            low*=2;high=2*high+1
        trace.append({'symbol':s,'freq':dict(model.freq),'before':old,'narrowed':narrowed,'ops':ops,'after':(low,high),'pending':pending,'emitted_delta':''.join(out[out_start:]),'emitted_bits':len(out)})
        model.update(s)
    pending+=1;emit(0 if low<Q else 1)
    return ''.join(out),trace

def arithmetic_decode(bits,freq,n,w=8,adaptive=False,threshold=None,enumerate_intervals=False):
    # bits is the logical coded stream followed by the mandatory w zero guard bits.
    need(type(w) is int and w in (8,16),'word width')
    need(type(n) is int and n>=0,'nonnegative symbol bound')
    need(all(b in '01' for b in bits),'not bits')
    Q=1<<(w-2);H=2*Q;TOP=4*Q-1;model=Model(freq,adaptive,threshold)
    need(sum(freq.values())<=Q-1,'frequency total')
    if adaptive:need(threshold<=Q-1,'frequency threshold')
    need(len(bits)>=w,'missing initial fill');value=int(bits[:w],2);pos=w;low=0;high=TOP;out=bytearray();trace=[]
    for _ in range(n+1):
        ranges,total=model.ranges();R=high-low+1
        if enumerate_intervals:
            # Independent inverse: enumerate disjoint integer subintervals, no scaled-value formula.
            choices=[(s,a,b) for s,a,b in ranges if low+R*a//total<=value<=low+R*b//total-1]
        else:
            scaled=((value-low+1)*total-1)//R
            choices=[(s,a,b) for s,a,b in ranges if a<=scaled<b]
        need(len(choices)==1,'invalid arithmetic state');s,a,b=choices[0];before=(low,high,value)
        high=low+R*b//total-1;low=low+R*a//total
        trace.append({'symbol':s,'before':before,'narrowed':(low,high),'bits_read_before_renormalize':pos})
        if s==EOF:
            need(len(out)==n,'early EOF');return bytes(out),trace,pos
        need(len(out)<n,'too many symbols');out.append(s)
        while True:
            if high<H:pass
            elif low>=H:low-=H;high-=H;value-=H
            elif low>=Q and high<3*Q:low-=Q;high-=Q;value-=Q
            else:break
            need(pos<len(bits),'truncated arithmetic stream')
            low*=2;high=2*high+1;value=2*value+int(bits[pos]);pos+=1
        model.update(s)
    raise Invalid('missing EOF')

def arc_encode(data,freq=None,w=8):
    freq=dict(sorted((freq or (Counter(data)|{EOF:1})).items()))
    bits,trace=arithmetic_encode(data,freq,w);payload,pad=pack_bits(bits+'0'*w)
    header=b'ARC1'+bytes([w])+struct.pack('>HII',len(freq),len(data),len(bits))
    header+=b''.join(struct.pack('>HH',s,f) for s,f in freq.items())
    return header+payload,{'freq':freq,'bits':bits,'guard_bits':w,'padding_bits':pad,'header_bytes':len(header),'trace':trace}

def arc_decode(blob,max_output=100000,max_bits=1000000,reference=False):
    w,n,count,freq,payload=parse_header(blob,b'ARC1',max_output,max_bits)
    need(count>=2,'missing arithmetic flush')
    # Includes guard before expansion, so physical-size check cannot be bypassed by padding.
    bits=unpack_bits(payload,count+w);need(bits[count:]=='0'*w,'guard bits')
    result,trace,read=arithmetic_decode(bits,freq,n,w,enumerate_intervals=reference)
    canonical_bits,_=arithmetic_encode(result,freq,w)
    need(canonical_bits==bits[:count],'noncanonical or incomplete arithmetic termination')
    return result,trace,read

def exact_interval(data,freq):
    model=Model(freq);lo=Fraction(0);hi=Fraction(1);trace=[]
    for s in list(data)+[EOF]:
        ranges,total=model.ranges();a,b=next((a,b) for x,a,b in ranges if x==s);R=hi-lo
        hi=lo+R*Fraction(b,total);lo=lo+R*Fraction(a,total)
        trace.append([s,str(lo),str(hi)])
    # Enumerate dyadic cells until a whole cell fits, not just its left endpoint.
    k=0
    while True:
        k+=1;scale=1<<k;j=(lo.numerator*scale+lo.denominator-1)//lo.denominator
        if Fraction(j+1,scale)<=hi:return lo,hi,f'{j:0{k}b}',trace

def lz77_encode(data,window=8,max_match=18):
    need(window>0 and max_match>=3,'LZ parameters');tokens=[];pos=0
    while pos<len(data):
        best=(0,0)
        for d in range(1,min(window,pos)+1):
            length=0
            while length<max_match and pos+length<len(data) and data[pos+length]==data[pos+length-d]:length+=1
            if length>best[0]:best=(length,d)
        if best[0]>=3:tokens.append(('M',*best));pos+=best[0]
        else:tokens.append(('L',data[pos]));pos+=1
    return tokens

def lz77_decode(tokens,window=8,max_output=100000,max_work=200000):
    need(type(window) is int and window>=1,'positive integer window')
    need(all(type(v) is int and v>=0 for v in (max_output,max_work)),'nonnegative integer budgets')
    out=bytearray();trace=[];work=0
    for token in tokens:
        need(isinstance(token,(tuple,list)),'LZ token shape')
        work+=1;need(work<=max_work,'token work budget')
        if len(token)==2 and token[0]=='L':
            s=token[1];need(type(s) is int and 0<=s<256,'literal')
            need(len(out)<max_output and work+1<=max_work,'literal budget');out.append(s);work+=1
        elif len(token)==3 and token[0]=='M':
            _,length,distance=token
            need(type(length) is int and type(distance) is int,'match integers')
            need(length>=3 and 1<=distance<=min(window,len(out)),'invalid backreference')
            # Subtraction form avoids fixed-width overflow; checked BEFORE copying/allocating.
            need(length<=max_output-len(out),'output budget')
            need(length<=max_work-work,'copy work budget')
            for _ in range(length):
                source=len(out)-distance;out.append(out[source]);work+=1
                trace.append({'source':source,'destination':len(out)-1,'byte':out[-1]})
        else:raise Invalid('token shape')
    return bytes(out),trace,work

def lz77_reference(tokens):
    out=b''
    for t in tokens:
        if t[0]=='L':out+=bytes([t[1]])
        else:
            _,length,distance=t;seed=out[-distance:]
            out+=(seed*((length+distance-1)//distance))[:length]
    return out

def lz78_encode(data):
    dictionary={b'':0};tokens=[];word=b''
    for s in data:
        candidate=word+bytes([s])
        if candidate in dictionary:word=candidate
        else:
            tokens.append((dictionary[word],s));dictionary[candidate]=len(dictionary);word=b''
    tokens.append((dictionary[word],None)) # explicit terminal reference, even for empty tail
    return tokens

def lz78_decode(tokens,max_output=100000,max_entries=10000,max_work=200000):
    # Parent-pointer dictionary: no duplicated long strings; validate length before materialization.
    need(all(type(v) is int and v>=0 for v in (max_output,max_work)),'nonnegative integer budgets')
    need(type(max_entries) is int and max_entries>=1,'dictionary budget includes empty root')
    dictionary=[(-1,None,0)];out=bytearray();work=0
    for at,token in enumerate(tokens):
        need(isinstance(token,(tuple,list)) and len(token)==2,'dictionary token shape')
        index,s=token
        work+=1;need(work<=max_work,'token work budget')
        need(type(index) is int and 0<=index<len(dictionary),'forward dictionary reference')
        terminal=s is None
        if terminal:need(at==len(tokens)-1,'nonfinal terminal')
        else:need(type(s) is int and 0<=s<256,'literal')
        length=dictionary[index][2]+(0 if terminal else 1)
        need(length<=max_output-len(out),'output budget');need(length<=max_work-work,'phrase work budget')
        if not terminal:need(len(dictionary)<max_entries,'dictionary budget')
        stack=[];cur=index
        while cur:
            parent,letter,_=dictionary[cur];stack.append(letter);cur=parent;work+=1
        out.extend(reversed(stack))
        if not terminal:out.append(s);work+=1;dictionary.append((index,s,length))
        else:return bytes(out)
    raise Invalid('missing terminal token')

def run():
    checks=[]
    def group(name,fn):fn();checks.append(name)
    data=b'ABABA';hf,hm=huf_encode(data);af,am=arc_encode(data)
    hd,ht=huf_decode(hf);ad,at,read=arc_decode(af)
    equal(hd,data,'Huffman endpoint');equal(ad,data,'arithmetic endpoint')
    lo,hi,ideal,idealtrace=exact_interval(data,am['freq'])
    lzt=lz77_encode(b'ABABABABA');lzo,lztrace,lzwork=lz77_decode(lzt)
    equal(lzo,b'ABABABABA','overlapping match')
    counts={'messages':0,'adaptive_messages':0,'length_tables':0,'malformed':0}
    def roundtrips():
        for n in range(9):
            for chars in product(b'AB',repeat=n):
                x=bytes(chars);h,meta=huf_encode(x);a,ar=arc_encode(x)
                equal(huf_decode(h)[0],x,'Huffman exhaustive')
                equal(huf_decode_reference(meta['bits'],meta['lengths']),x,'tree reference')
                equal(arc_decode(a)[0],x,'arithmetic exhaustive')
                equal(arc_decode(a,reference=True)[0],x,'enumerated interval reference')
                equal(lz77_decode(lz77_encode(x))[0],x,'LZ77 exhaustive')
                equal(lz77_reference(lz77_encode(x)),x,'periodic reference')
                equal(lz78_decode(lz78_encode(x)),x,'LZ78 exhaustive');counts['messages']+=1
    group('511 binary strings: both complete formats plus LZ77/LZ78',roundtrips)
    def alphabets():
        for x,w in [(bytes(range(256)),16),(b'\x00\xff\x00',8),(b'A'*1024,16)]:
            h,hm=huf_encode(x);a,am=arc_encode(x,w=w)
            equal(huf_decode(h)[0],x,'full-byte Huffman');equal(arc_decode(a)[0],x,'full-byte arithmetic')
    group('full 256-byte alphabet, NUL/FF and long single-letter data',alphabets)
    def break_even():
        for n in range(1,65):equal(len(huf_encode(b'A'*n)[0]),20+(n+8)//8,'single-letter complete size')
        equal(next(n for n in range(1,65) if len(huf_encode(b'A'*n)[0])<n),25,'first strict saving')
        for n in (23,24):equal(len(huf_encode(b'A'*n)[0]),n,'break-even equality')
    group('single-letter HUF1 complete cost and first strict saving n=25',break_even)
    def tables():
        for n in range(1,5):
            for lens in product(range(1,5),repeat=n):
                d=dict(enumerate(lens));kraft=sum(Fraction(1,1<<l) for l in lens)
                if kraft<=1:equal(canonical(d),canonical_reference(d),'enumerated canonical table')
                else:rejected(lambda:canonical(d))
                counts['length_tables']+=1
    group('340 length tables against Kraft and independent leaf enumeration',tables)
    def adaptive():
        for n in range(8):
            for chars in product(b'AB',repeat=n):
                x=bytes(chars);freq={65:1,66:1,EOF:1};bits,tr=arithmetic_encode(x,freq,8,True,7)
                equal(arithmetic_decode(bits+'0'*8,freq,len(x),8,True,7)[0],x,'adaptive inverse')
                equal(arithmetic_decode(bits+'0'*8,freq,len(x),8,True,7,True)[0],x,'adaptive independent inverse');counts['adaptive_messages']+=1
        x=b'A'*1000+b'B'*1000;freq={65:1,66:1,EOF:1};bits,tr=arithmetic_encode(x,freq,8,True,7)
        equal(arithmetic_decode(bits+'0'*8,freq,len(x),8,True,7)[0],x,'repeated rescaling')
    group('255 adaptive messages plus 2000 symbols with repeated frequency rescaling',adaptive)
    def bad(fn):rejected(fn);counts['malformed']+=1
    def malformed():
        for blob,decoder in [(hf,huf_decode),(af,arc_decode)]:
            for end in range(len(blob)):bad(lambda b=blob[:end],d=decoder:d(b))
            bad(lambda b=blob,d=decoder:d(b+b'\0'))
            bad(lambda b=blob,d=decoder:d(b,max_output=4))
        for ls in ({0:0},{0:16},{0:1,1:1,2:1},{}):bad(lambda ls=ls:canonical(ls))
        bad(lambda:arithmetic_encode(b'A',{65:63,EOF:1}))
        bad(lambda:arithmetic_encode(b'A',{65:0,EOF:1}))
        for tokens in [[('M',3,1)],[('L',65),('M',3,0)],[('L',65),('M',3,2)],[('L',65),('M',10**30,1)]]:bad(lambda t=tokens:lz77_decode(t,max_output=100,max_work=100))
        bad(lambda:lz77_decode([('L',65),('M',10,1)],max_work=5))
        for key in ('max_output','max_work'):
            for v in (-1,True,1.5):bad(lambda key=key,v=v:lz77_decode([],**{key:v}))
        bad(lambda:lz77_decode([],window=0))
        for key in ('max_output','max_work'):
            for v in (-1,True,1.5):bad(lambda key=key,v=v:lz78_decode([(0,None)],**{key:v}))
        for v in (0,-1,True,1.5):bad(lambda v=v:lz78_decode([(0,None)],max_entries=v))
        for d,b in [(huf_decode,hf),(arc_decode,af)]:
            for key in ('max_output','max_bits'):
                for v in (-1,True,1.5):bad(lambda d=d,b=b,key=key,v=v:d(b,**{key:v}))
        equal(lz77_decode([],max_output=0,max_work=0)[0],b'','empty zero budgets')
        equal(lz78_decode([(0,None)],max_output=0,max_work=1,max_entries=1),b'','empty dictionary initial capacity')
        bad(lambda:lz78_decode([(1,65),(0,None)]))
        bad(lambda:lz78_decode([(0,65)]))
        bad(lambda:lz78_decode([(0,None),(0,65)]))
        bad(lambda:lz78_decode(lz78_encode(b'ABABABA'),max_entries=2))
        bad(lambda:lz78_decode(lz78_encode(b'ABABABA'),max_output=4))
        # Raw mutated headers: huge count is rejected before table/payload expansion.
        q=bytearray(hf);q[6:10]=(2**32-1).to_bytes(4,'big');bad(lambda:huf_decode(bytes(q)))
        q=bytearray(hf);q[17:19]=q[14:16];bad(lambda:huf_decode(bytes(q)))
        bad(lambda:huf_decode(bytes.fromhex('485546310001000000000000000201000200')))
        q=bytearray(hf);q[20:22]=(255).to_bytes(2,'big');bad(lambda:huf_decode(bytes(q)))
        q=bytearray(hf);q[-1]|=1;bad(lambda:huf_decode(bytes(q)))
        q=bytearray(af);q[-1]|=1;bad(lambda:arc_decode(bytes(q)))
        q=bytearray(af);q[-2]|=1;bad(lambda:arc_decode(bytes(q)))
    group('truncation, trailing data, table, guard/pad, frequency and resource failures',malformed)
    adaptive_bits,adaptive_trace=arithmetic_encode(b'ABABA',{65:1,66:1,EOF:1},8,True,7)
    return {'status':'passed','teaching_format':True,'input_hex':data.hex(),'input_bits':8*len(data),'huffman':hm|{'file_hex':hf.hex(),'file_bits':len(hf)*8,'decode_trace':ht},'arithmetic':am|{'file_hex':af.hex(),'file_bits':len(af)*8,'decode_trace':at,'decode_bits_read':read},'ideal_arithmetic':{'interval':[str(lo),str(hi)],'width':str(hi-lo),'dyadic_cell_bits':ideal,'trace':idealtrace},'adaptive':{'bits':adaptive_bits,'trace':adaptive_trace,'rule':'encode/decode, increment non-EOF, if total>=7 ceil-halving'},'lz77':{'tokens':lzt,'output':lzo.decode(),'copy_trace':lztrace,'work':lzwork},'lz78':{'ABABA':lz78_encode(b'ABABA'),'ABABABA':lz78_encode(b'ABABABA')},'test_counts':counts,'checks':checks,'independence':'Author cross-check: leaf enumeration and tree decoder; arithmetic inverse enumerates integer subintervals; LZ77 periodic-string construction. These are not an external review or an arbitrary-input proof.'}
if __name__=='__main__':print(json.dumps(run(),ensure_ascii=False,indent=2))
