#!/usr/bin/env python3
"""DB-4 finite teaching model. Python standard library only; no database I/O.
Counts simulate full-page transfers under the published eager schedules.
This is a finite checker, not a proof of arbitrary SQL or storage correctness.
"""
from collections import Counter, defaultdict
from itertools import permutations, product
from math import ceil
from fractions import Fraction
import json

R = [(1,'A'),(2,'B'),(1,'A'),(3,'C'),(2,'D'),(4,'E')]
S = [(2,'u'),(1,'v'),(2,'u'),(3,'w'),(1,'x'),(5,'y'),(4,'z'),(2,'t')]
CAPACITY, FRAMES = 2, 4

def pages(rows):
    return ceil(len(rows) / CAPACITY)

def oracle(r, s):
    return Counter((a[0], a[1], b[1]) for a in r for b in s if a[0] == b[0])

def bnlj(r, s):
    out = Counter()
    blocks = [r[i:i+(FRAMES-2)*CAPACITY] for i in range(0,len(r),(FRAMES-2)*CAPACITY)]
    for block in blocks:
        for b in s:
            for a in block:
                if a[0] == b[0]: out[(a[0],a[1],b[1])] += 1
    # Eager model reads every specified outer page and each full inner scan.
    return out, pages(r) + len(blocks)*pages(s)

def grace(r, s):
    rp = [[a for a in r if a[0]%2 == k] for k in range(2)]
    sp = [[b for b in s if b[0]%2 == k] for k in range(2)]
    temporary_pages = sum(pages(x) for x in rp+sp)
    io = pages(r)+pages(s)+temporary_pages
    out, fallbacks = Counter(), 0
    for a,b in zip(rp,sp):
        if len(a) > (FRAMES-2)*CAPACITY:
            bag, cost = bnlj(a,b)
            out.update(bag);io += cost;fallbacks += 1
        else:
            table = defaultdict(list)
            for item in a: table[item[0]].append(item)
            for item in b:
                for match in table[item[0]]:
                    out[(item[0], match[1], item[1])] += 1
            io += pages(a)+pages(b)
    return out,io,temporary_pages,fallbacks

def merge(r, s):
    a,b = sorted(r),sorted(s)
    i=j=0;out=Counter();max_rgroup=0
    while i<len(a) and j<len(b):
        if a[i][0]<b[j][0]: i+=1
        elif b[j][0]<a[i][0]: j+=1
        else:
            k=a[i][0];u=i;v=j
            while i<len(a) and a[i][0]==k:i+=1
            while j<len(b) and b[j][0]==k:j+=1
            max_rgroup=max(max_rgroup,i-u)
            # Groups are crossed completely; same-valued rows stay separate.
            for aa in a[u:i]:
                for bb in b[v:j]:out[(k,aa[1],bb[1])]+=1
    return out,max_rgroup

def dp_example():
    tables={'R':R,'S':S,'T':[(1,'t1'),(2,'t2')]}
    def join_keys(a,b):return [x for x in a for y in b if x==y]
    dp={frozenset([t]):0 for t in tables}
    counts={frozenset([t]):len(rows) for t,rows in tables.items()}
    for a,b in [('R','S'),('R','T'),('S','T')]:
        counts[frozenset([a,b])]=len(join_keys([x[0] for x in tables[a]],[x[0] for x in tables[b]]))
        dp[frozenset([a,b])]=len(tables[a])*len(tables[b])
    last={j:dp[frozenset(tables)-{j}]+counts[frozenset(tables)-{j}]*len(tables[j]) for j in tables}
    assert last=={'R':46,'S':44,'T':72}
    brute=[]
    for a,b,c in permutations(tables):
        n=counts[frozenset([a,b])]
        brute.append((a+b+c,len(tables[a])*len(tables[b])+n*len(tables[c])))
    assert min(x[1] for x in brute)==44
    return {'last_relation_costs':last,'minimum':44,'all_six_orders':brute}

def page_move():
    p=bytearray(128)
    p[104:128]=b'A'*24;p[88:104]=b'B'*16;p[64:88]=b'C'*24
    # Snapshot copy deliberately has memmove semantics for overlapping ranges.
    p[80:104]=bytes(p[64:88]);p[40:80]=b'D'*40
    assert p[80:104]==b'C'*24 and p[104:128]==b'A'*24
    assert 128-(16+4*4+24+24+40)==8
    return {'source':[64,88],'target':[80,104],'overlap_bytes':8,'free_after_insert':8}

def recovery():
    # Stable-prefix physical redo, then undo any transaction without stable COMMIT.
    log=[{'lsn':40,'tx':'T','kind':'UPDATE','before':7,'after':9},
         {'lsn':50,'tx':'T','kind':'COMMIT'}]
    def restart(stable_lsn,disk_value,disk_lsn):
        stable=[r for r in log if r['lsn']<=stable_lsn]
        winners={r['tx'] for r in stable if r['kind']=='COMMIT'}
        value,lsn=disk_value,disk_lsn
        actions=[]
        for r in stable:
            if r['kind']=='UPDATE' and lsn<r['lsn']:
                value,lsn=r['after'],r['lsn'];actions.append('redo '+str(r['lsn']))
        for r in reversed(stable):
            if r['kind']=='UPDATE' and r['tx'] not in winners:
                value=r['before'];actions.append('undo '+str(r['lsn']))
        return value,actions
    loser,loser_actions=restart(40,9,40)
    winner,winner_actions=restart(50,7,10)
    assert loser==7 and winner==9
    assert loser_actions==['undo 40'] and winner_actions==['redo 40']
    assert not (40>=60)  # Forbidden flush: pageLSN60, flushedLSN40.
    return {'loser_after_restart':loser,'loser_actions':loser_actions,
            'winner_after_restart':winner,'winner_actions':winner_actions,
            'flush_page60_after_log40':'rejected'}

def external_sort(page_count,frames):
    runs=[min(frames,page_count-i) for i in range(0,page_count,frames)]
    trace=[runs[:]];io=2*page_count
    while len(runs)>1:
        runs=[sum(runs[i:i+frames-1]) for i in range(0,len(runs),frames-1)]
        trace.append(runs[:]);io+=2*page_count
    # Full-copy schedule includes copying singleton runs on every pass.
    return io,trace

def aggregate_cost(rows):
    parts=[[r for r in rows if r[0] in keys] for keys in [{'a','c'},{'b','d'}]]
    temporary=sum(pages(part) for part in parts)
    states={}
    for part in parts:
        local=defaultdict(lambda:[0,0])
        for k,v in part:local[k][0]+=1;local[k][1]+=v
        assert pages(local)<=1  # Each state is a 64B slot in this small example.
        states.update(local)
    io=pages(rows)+2*temporary+pages(states)
    return states,io,temporary

def main():
    expected=oracle(R,S)
    b,bi=bnlj(R,S);h,hi,tmp,fb=grace(R,S);m,g=merge(R,S)
    assert b==h==m==expected and len(expected)==8 and sum(expected.values())==12
    assert (bi,hi,tmp,fb,g)==(11,23,8,0,2)
    assert bnlj(S,R)[1]==10
    sort_io=2*(pages(R)+pages(S));merge_io=pages(R)+pages(S)
    assert sort_io+merge_io==21
    # Exhaust every sequence up to length 3 over three records, including exact duplicates.
    universe=[(1,'a'),(1,'b'),(2,'a')]
    samples=[list(t) for n in range(4) for t in product(universe,repeat=n)]
    checked=0
    for rr in samples:
        for ss in samples:
            truth=oracle(rr,ss)
            assert bnlj(rr,ss)[0]==grace(rr,ss)[0]==merge(rr,ss)[0]==truth
            checked+=1
    hot_r=[(1,'a')]*8;hot_s=[(1,'x')]*7
    gh=grace(hot_r,hot_s)
    assert gh[0]==oracle(hot_r,hot_s) and gh[3]==1 and sum(gh[0].values())==56
    x=[1,1,1,0,1,0]
    estimate=6*Fraction(sum(x),6)**2
    assert estimate==Fraction(8,3)
    assert Fraction(6*8,5)==Fraction(48,5)
    assert estimate*8/5==Fraction(64,15)
    assert ceil(estimate/2)==ceil(4/2)==2
    rows=[('a',1),('b',2),('a',3),('c',4),('b',5),('d',6)]
    states,aggregate_io,aggregate_temp=aggregate_cost(rows)
    assert aggregate_io==13 and aggregate_temp==4
    external_io,run_trace=external_sort(9,FRAMES)
    fused_io=external_io-9  # Final write omitted, downstream output counted separately.
    assert external_io==36 and fused_io==27
    external13,trace13=external_sort(13,FRAMES)
    assert external13==78 and trace13==[[4,4,4,1],[12,1],[13]]
    assert dict(states)=={'a':[2,4],'b':[2,7],'c':[1,4],'d':[1,6]}
    latest={}
    versions=[('a',1,'PUT',10),('b',2,'PUT',20),('a',3,'PUT',11),('b',4,'DEL',None),('c',5,'PUT',30)]
    for k,seq,kind,val in versions:
        if k not in latest or latest[k][0]<seq:latest[k]=(seq,kind,val)
    assert {k:v[2] for k,v in latest.items() if v[1]=='PUT'}=={'a':11,'c':30}
    lsm_output=[v for v in latest.values() if v[1]=='PUT']
    lsm_io=pages(versions[:2])+pages(versions[2:])+pages(lsm_output)
    assert lsm_io==4
    frame_stages={'bnlj':[2,1,1],'grace_partition':[1,2],'grace_join':[2,1,1],'merge':[1,1,1,1]}
    assert all(sum(parts)<=FRAMES for parts in frame_stages.values())
    result={'status':'passed','bag':[{'tuple':list(k),'multiplicity':v} for k,v in sorted(expected.items())],
        'io_excluding_output':{'bnlj_r_outer':bi,'bnlj_s_outer':10,'grace':hi,'sort_merge':sort_io+merge_io},
        'output_pages':12,'total_io':{'bnlj_r_outer':23,'bnlj_s_outer':22,'grace':35,'sort_merge':33},
        'exhaustive_bag_pairs':checked,'hot_key_output':56,'hash_temporary_pages':tmp,
        'correlation_estimate':str(estimate),'true_selected_rows':4,'ndv_join_estimate':'48/5','true_join_rows':12,
        'dp':dp_example(),'slot_page':page_move(),'recovery':recovery(),
        'aggregate':dict(states),'aggregation_io':aggregate_io,'external_sort_io':external_io,'fused_sort_io':fused_io,
        'sort_run_trace':run_trace,'sort13_io':external13,'sort13_run_trace':trace13,
        'declared_frame_stages':frame_stages,
        'memory_scope':'Host Python lists/sort are not a bounded-buffer implementation; frame sums check the published schedule only.',
        'lsm_latest':{'a':11,'c':30},'lsm_data_compaction_io':lsm_io}
    print(json.dumps(result,ensure_ascii=False,indent=2))
if __name__=='__main__':main()
