#!/usr/bin/env python3
"""Finite teaching model: sloppy placement, durable hinted handoff and snapshot
Merkle range reconciliation. Standard library; stdout only. No actual network or
fsync. Fixed honest identities/epoch; atomic stable steps survive crash(). Input
records obey per-key unique LWW tags; tag generation is outside this interface.
SHA-256 is a conditional collision-resistant tool, not a collision-free oracle.
"""
from dataclasses import dataclass,asdict,replace
from hashlib import sha256
from copy import deepcopy
from itertools import combinations,product
import json,random,struct

def need(ok,msg):
    if not ok:raise ValueError(msg)
def natural(x):return type(x)is int and x>=0
def reject(call):
    try:call()
    except ValueError as e:return str(e)
    raise RuntimeError('required rejection missing')

@dataclass(frozen=True)
class Record:
    key:int
    tag:int
    deleted:bool
    payload:bytes
    def __post_init__(self):
        need(natural(self.key)and self.key<2**32,'u32 key')
        need(natural(self.tag)and self.tag<2**64,'u64 per-key tag')
        need(type(self.deleted)is bool and type(self.payload)is bytes and len(self.payload)<2**32,'record value')
        need(not self.deleted or self.payload==b'','canonical tombstone payload')
    def wire(self):return struct.pack('>IQBI',self.key,self.tag,self.deleted,len(self.payload))+self.payload

def join(x,y):
    if x is None:return y
    need(x.key==y.key,'same register')
    if x.tag==y.tag:need(x==y,'same tag must identify same contents')
    return x if x.tag>=y.tag else y

def placement(preference,N,reachable):
    need(natural(N)and 1<=N<=len(preference),'replication count')
    need(len(set(preference))==len(preference),'distinct physical preferences')
    homes=tuple(preference[:N]);reachable=set(reachable)
    selected=[p for p in preference if p in reachable][:N]
    selected_set=set(selected);home_set=set(homes)
    direct=[(h,h)for h in homes if h in selected_set]
    missed=[h for h in homes if h not in selected_set]
    backups=[p for p in selected if p not in home_set]
    return direct+list(zip(missed,backups))  # (logical home, actual physical recipient)

@dataclass(frozen=True)
class Put:
    epoch:int
    operation:str
    home:str
    record:Record
@dataclass(frozen=True)
class PutAck:
    request:Put
    physical:str
    durable_tag:int
@dataclass(frozen=True)
class Transfer:
    epoch:int
    source:str
    home:str
    record:Record
@dataclass(frozen=True)
class TransferAck:
    request:Transfer
    receiver:str
    durable_tag:int

class Node:
    def __init__(self,name,members,homes,epoch=1):
        need(len(set(members))==len(members)and len(set(homes))==len(homes)and name in members and set(homes)<=set(members),'registered identities')
        need(natural(epoch),'epoch');self.name=name;self.members=tuple(members);self.homes=tuple(homes);self.epoch=epoch
        self.data={};self.hints={};self.revision=0 # all stable state, including revision
        self.pending={}                         # volatile reception only
    def stage_put(self,req):
        need(isinstance(req,Put)and req.epoch==self.epoch and req.home in self.homes,'put identity/epoch/home')
        need(self.name==req.home or self.name not in self.homes,'backup outside fixed home set')
        token=('PUT',req.operation)
        if token in self.pending:need(self.pending[token]==req,'operation content reuse')
        self.pending[token]=req;return token
    def flush_put(self,token):
        need(token in self.pending and isinstance(self.pending[token],Put),'staged put')
        req=self.pending[token]
        if self.name==req.home:
            before=self.data.get(req.record.key);after=join(before,req.record);self.data[req.record.key]=after
        else:
            slot=(req.home,req.record.key);before=self.hints.get(slot);after=join(before,req.record);self.hints[slot]=after
        if before!=after:self.revision+=1
        del self.pending[token]
        return PutAck(req,self.name,after.tag) # only after the atomic stable join
    def start_transfer(self,home,key):
        need((home,key)in self.hints,'retained hint');return Transfer(self.epoch,self.name,home,self.hints[home,key])
    def stage_transfer(self,req):
        need(isinstance(req,Transfer)and req.epoch==self.epoch and req.home==self.name and self.name in self.homes and req.source in self.members and req.source not in self.homes,'handoff endpoint/epoch')
        token=('TRANSFER',req);self.pending[token]=req;return token
    def flush_transfer(self,token):
        need(token in self.pending and isinstance(self.pending[token],Transfer),'staged transfer')
        req=self.pending[token];before=self.data.get(req.record.key);after=join(before,req.record)
        self.data[req.record.key]=after
        if before!=after:self.revision+=1
        del self.pending[token]
        return TransferAck(req,self.name,after.tag)
    def finish_transfer(self,ack):
        need(isinstance(ack,TransferAck),'handoff ACK type');req=ack.request
        need(req.epoch==self.epoch and req.source==self.name and ack.receiver==req.home and req.home in self.homes and ack.durable_tag>=req.record.tag,'matching durable ACK')
        slot=(req.home,req.record.key);current=self.hints.get(slot)
        if current!=req.record:return 'STALE_OR_DUPLICATE_ACK'
        del self.hints[slot];self.revision+=1;return 'RETIRED' # atomic stable deletion
    def crash(self):self.pending.clear() # stable data/hints/revision survive
    def absorb(self,records):
        # Caller has already validated the whole repair payload. This model uses
        # one atomic stable batch; copy cost is accounted separately in the page.
        updates={}
        for rec in records:updates[rec.key]=join(updates.get(rec.key,self.data.get(rec.key)),rec)
        for k,v in updates.items():
            if self.data.get(k)!=v:self.data[k]=v;self.revision+=1
    def show(self):return dict(node=self.name,revision=self.revision,data=[self.data[k]for k in sorted(self.data)],hints=[dict(home=h,record=self.hints[h,k])for h,k in sorted(self.hints)],volatile=len(self.pending))

class Write:
    def __init__(self,epoch,operation,record,targets,W):
        need(type(operation)is str and operation and natural(W)and W>=1,'write identity/quorum')
        need(len({p for h,p in targets})==len(targets)and len({h for h,p in targets})==len(targets),'distinct home/physical assignment')
        self.requests={p:Put(epoch,operation,h,record)for h,p in targets};self.W=W;self.acknowledged=set()
    def accept(self,ack):
        need(isinstance(ack,PutAck)and ack.physical in self.requests and ack.request==self.requests[ack.physical]and ack.durable_tag>=ack.request.record.tag,'matching write ACK')
        self.acknowledged.add(ack.physical);return len(self.acknowledged)>=self.W
    def status(self):return dict(complete=len(self.acknowledged)>=self.W,physical=sorted(self.acknowledged),needed=self.W)

def read_responses(operation,key,replies,R,selected):
    need(type(operation)is str and operation and natural(R)and R>=1,'read identity/quorum')
    allowed=set(selected);need(len(allowed)==len(selected),'distinct selected read responders');seen={}
    for op,physical,record in replies:
        need(op==operation and physical in allowed and (record is None or isinstance(record,Record)and record.key==key),'matching read reply')
        if physical in seen:need(seen[physical]==record,'one cached response per request and physical node')
        seen[physical]=record
    if len(seen)<R:return dict(status='INCOMPLETE',physical=sorted(seen),record=None)
    winner=None
    for r in seen.values():
        if r is not None:winner=join(winner,r)
    return dict(status='ABSENT'if winner is None else'VALUE',physical=sorted(seen),record=winner)

@dataclass(frozen=True)
class Config:
    epoch:int=1
    low:int=0
    key_bits:int=4
    depth:int=3
    schema:str='LWW-u32-u64-SHA256-v1'
    def __post_init__(self):
        need(natural(self.epoch)and natural(self.low)and self.low<2**32,'range epoch/start')
        need(natural(self.key_bits)and self.key_bits<=12 and natural(self.depth)and self.depth<=self.key_bits,'teaching tree width/depth 0..12')
        need(self.low+2**self.key_bits<=2**32 and self.schema=='LWW-u32-u64-SHA256-v1','range/schema')
    def bucket(self,key):
        need(self.low<=key<self.low+2**self.key_bits,'key inside range');return (key-self.low)//(2**(self.key_bits-self.depth))

def bucket_wire(config,index,records):
    need(natural(index)and index<2**config.depth,'bucket position');previous=-1;chunks=[struct.pack('>I',len(records))]
    for r in records:
        need(isinstance(r,Record)and config.bucket(r.key)==index and r.key>previous,'sorted unique complete bucket records')
        chunks.append(r.wire());previous=r.key
    return b''.join(chunks)
def leaf(index,payload):return sha256(b'\x00'+struct.pack('>II',index,len(payload))+payload).digest()
def parent(level,left,right):
    need(natural(level)and 1<=level<=12 and type(left)is bytes and type(right)is bytes and len(left)==len(right)==32,'digest level/width');return sha256(b'\x01'+bytes([level])+left+right).digest()
@dataclass(frozen=True)
class Children:
    snapshot:tuple
    level:int
    index:int
    digests:tuple
@dataclass(frozen=True)
class Bucket:
    snapshot:tuple
    index:int
    records:tuple

class Snapshot:
    def __init__(self,node,config):
        need(node.name in node.homes and node.epoch==config.epoch,'snapshot owner/config')
        self.config=config;self.identity=(node.name,node.epoch,node.revision,config.low,config.key_bits,config.depth,config.schema)
        groups=[[]for _ in range(2**config.depth)]
        # Out-of-range records are outside this repair session, not discarded.
        for key in sorted(node.data):
            if config.low<=key<config.low+2**config.key_bits:groups[config.bucket(key)].append(node.data[key])
        self.records=tuple(tuple(g)for g in groups)
        wire=[bucket_wire(config,i,g)for i,g in enumerate(self.records)];self.encoded_bytes=sum(map(len,wire))
        levels=[tuple(leaf(i,v)for i,v in enumerate(wire))]
        for j in range(1,config.depth+1):levels.append(tuple(parent(j,levels[-1][2*i],levels[-1][2*i+1])for i in range(len(levels[-1])//2)))
        self.levels=tuple(levels);self.root=self.levels[-1][0]
    def children(self,level,index):
        need(natural(level)and 1<=level<=self.config.depth and natural(index)and index<len(self.levels[level]),'internal tree coordinate')
        return Children(self.identity,level,index,(self.levels[level-1][2*index],self.levels[level-1][2*index+1]))
    def bucket(self,index):
        need(natural(index)and index<len(self.records),'leaf index');return Bucket(self.identity,index,self.records[index])

def validate_children(reply,snapshot,level,index,expected):
    need(isinstance(reply,Children)and reply.snapshot==snapshot and (reply.level,reply.index)==(level,index)and len(reply.digests)==2,'child reply identity/coordinate')
    need(parent(level,*reply.digests)==expected,'children reconstruct accepted parent');return reply.digests

def validate_bucket(reply,identity,config,index,digest):
    need(isinstance(reply,Bucket)and reply.snapshot==identity and reply.index==index,'bucket reply identity/position')
    raw=bucket_wire(config,index,reply.records);need(leaf(index,raw)==digest,'bucket contents reconstruct accepted leaf');return reply.records

def reconcile(left,right,left_node,right_node):
    need(left.config==right.config,'same range descriptor before root comparison')
    need(left.identity[0]==left_node.name and right.identity[0]==right_node.name and left.config.epoch==left_node.epoch==right_node.epoch,'snapshot session endpoints')
    config=left.config;work=[(config.depth,0,left.root,right.root)];visits=[];different=[];to_left=[];to_right=[];payload_bytes=0
    while work:
        j,i,x,y=work.pop();same=x==y;visits.append(dict(level=j,index=i,equal=same))
        if same:continue
        if j:
            lx=validate_children(left.children(j,i),left.identity,j,i,x);ry=validate_children(right.children(j,i),right.identity,j,i,y)
            work.append((j-1,2*i+1,lx[1],ry[1]));work.append((j-1,2*i,lx[0],ry[0]))
        else:
            l=validate_bucket(left.bucket(i),left.identity,config,i,x);r=validate_bucket(right.bucket(i),right.identity,config,i,y)
            different.append(i);to_left.extend(r);to_right.extend(l);payload_bytes+=len(bucket_wire(config,i,l))+len(bucket_wire(config,i,r))
    # Validate join compatibility before any apply. Different illegal payloads
    # under one per-key tag are outside the LWW contract and must be rejected.
    merged={}
    for rec in to_left+to_right:merged[rec.key]=join(merged.get(rec.key),rec)
    for node,rows in [(left_node,to_left),(right_node,to_right)]:
        for rec in rows:join(node.data.get(rec.key),rec)
    left_node.absorb(to_left);right_node.absorb(to_right)
    return dict(status='SNAPSHOTS_ABSORBED',left=left.identity,right=right.identity,compared=len(visits),remote_digest_bytes=32*len(visits),different_buckets=different,record_occurrences=len(to_left)+len(to_right),bucket_payload_bytes=payload_bytes,visits=visits)

def main_trace():
    members=tuple('ABCDEF');homes=tuple('ABC');nodes={p:Node(p,members,homes)for p in members};old=Record(3,0,False,b'old');v7=Record(3,7,False,b'v7');v9=Record(3,9,False,b'v9')
    for h in homes:nodes[h].absorb([old])
    chosen=placement(members,3,'DEF');w=Write(1,'client-write-7',v7,chosen,2)
    ackD=nodes['D'].flush_put(nodes['D'].stage_put(w.requests['D']));w.accept(ackD);w.accept(ackD);need(not w.status()['complete'],'duplicate physical ACK not enough')
    ackE=nodes['E'].flush_put(nodes['E'].stage_put(w.requests['E']));w.accept(ackE)
    read=read_responses('read-after-7',3,[('read-after-7',p,nodes[p].data[3])for p in 'AB'],2,homes)
    need(w.status()['complete']and read['record'].tag==0,'sloppy completed write / stale subsequent read')
    before={p:nodes[p].show()for p in members};D,A=nodes['D'],nodes['A'];req7=D.start_transfer('A',3)
    A.stage_transfer(req7);A.crash();need(A.data[3].tag==0 and D.hints['A',3].tag==7,'receive without durable commit')
    lost=A.flush_transfer(A.stage_transfer(req7));A.crash();need(A.data[3].tag==7 and D.hints['A',3].tag==7,'durable commit / lost ACK')
    retry=A.flush_transfer(A.stage_transfer(req7));need(retry==lost,'duplicate identical handoff')
    req9=Put(1,'client-write-9','A',v9);D.flush_put(D.stage_put(req9));stale=D.finish_transfer(retry);need(stale=='STALE_OR_DUPLICATE_ACK'and D.hints['A',3].tag==9,'old ACK cannot erase newer hint')
    packet9=D.start_transfer('A',3);ack9=A.flush_transfer(A.stage_transfer(packet9));retired=D.finish_transfer(ack9);D.crash();A.crash();need(retired=='RETIRED'and not D.hints and A.data[3].tag==9,'durable retirement')
    # A higher receiver version still acknowledges an older transfer's lower bound.
    D.flush_put(D.stage_put(Put(1,'retry-old-9','A',v9)));A.absorb([Record(3,11,False,b'v11')]);above=A.flush_transfer(A.stage_transfer(D.start_transfer('A',3)));D.finish_transfer(above);need(A.data[3].tag==11,'receiver maximum is not overwritten')
    def pair():
        a=Node('A',members,homes);c=Node('C',members,homes)
        a.absorb([Record(0,1,False,b'a'),Record(3,2,False,b'c'),Record(8,4,False,b'h'),Record(14,6,False,b'n')]);c.absorb([Record(0,1,False,b'a'),Record(3,1,False,b'old'),Record(9,3,False,b'i'),Record(14,7,True,b'')]);return a,c
    config=Config();a,c=pair();sa,sc=Snapshot(a,config),Snapshot(c,config);first=reconcile(sa,sc,a,c)
    need(first['compared']==13 and first['different_buckets']==[1,4,7]and first['record_occurrences']==6 and a.data==c.data,'three changed buckets')
    clean=dict(A=a.show(),C=c.show(),roots=[sa.root.hex(),sc.root.hex()],repair=first)
    a,c=pair();sa,sc=Snapshot(a,config),Snapshot(c,config);a.absorb([Record(3,9,False,b'new')]);concurrent=reconcile(sa,sc,a,c);middle=dict(A=a.show(),C=c.show())
    need(a.data[3].tag==9 and c.data[3].tag==2,'old snapshots do not imply current equality')
    next_round=reconcile(Snapshot(a,config),Snapshot(c,config),a,c);need(next_round['compared']==7 and a.data==c.data,'new snapshot carries new write')
    errors=[reject(lambda:validate_children(replace(sc.children(3,0),snapshot=('C',1,999)),sc.identity,3,0,sc.root)),reject(lambda:validate_bucket(sc.bucket(1),sc.identity,config,4,sc.levels[0][4])),reject(lambda:reconcile(sa,Snapshot(c,Config(depth=2)),a,c))]
    return dict(sloppy=dict(assignment=chosen,write=w.status(),read=read,before_handoff=before,old_ack=stale,retirement=retired,after_A=A.show(),after_D=D.show(),higher_receiver_ack=above),merkle=dict(no_concurrent_write=clean,with_concurrent_write=dict(first=concurrent,middle=middle,second=next_round,final_A=a.show(),final_C=c.show()),explicit_rejections=errors))

def tests():
    rng=random.Random(161616);counts=dict(placements=0,handoff_windows=0,snapshot_pairs=0,changed_buckets=0)
    members=tuple('ABCDEF');homes=tuple('ABC')
    for N in range(1,7):
        for mask in range(64):
            reach={p for i,p in enumerate(members)if mask>>i&1};picked=placement(members,N,reach);expected=[p for p in members if p in reach][:N]
            need({p for h,p in picked}==set(expected)and len({h for h,p in picked})==len(picked),'literal preference selection');counts['placements']+=1
    for early,late,receiver,crash_at in product(range(1,5),range(5,9),range(0,11),range(4)):
        D=Node('D',members,homes);A=Node('A',members,homes)
        r=Record(3,early,False,str(early).encode());new=Record(3,late,False,str(late).encode());A.absorb([Record(3,receiver,False,str(receiver).encode())]);D.flush_put(D.stage_put(Put(1,'old','A',r)));p=D.start_transfer('A',3)
        if crash_at==0:A.stage_transfer(p);A.crash()
        ack=A.flush_transfer(A.stage_transfer(p))
        if crash_at==1:A.crash()
        D.flush_put(D.stage_put(Put(1,'new','A',new)))
        if crash_at==2:D.crash()
        need(D.finish_transfer(ack)=='STALE_OR_DUPLICATE_ACK'and D.hints['A',3]==new,'stale handoff guard')
        ack2=A.flush_transfer(A.stage_transfer(D.start_transfer('A',3)));D.finish_transfer(ack2)
        if crash_at==3:D.crash();A.crash()
        need(not D.hints and A.data[3].tag==max(receiver,late),'final durable lower bound');counts['handoff_windows']+=1
    for trial in range(1000):
        bits=rng.randrange(0,7);depth=rng.randrange(bits+1);config=Config(key_bits=bits,depth=depth);nodes=[Node(p,members,homes)for p in 'AC']
        for node in nodes:
            rows=[]
            for key in range(2**bits):
                if rng.randrange(3):
                    tag=rng.randrange(5);deleted=tag==4;rows.append(Record(key,tag,deleted,b''if deleted else str(tag).encode()))
            node.absorb(rows)
        x,y=[Snapshot(node,config)for node in nodes];expected={}
        for node in nodes:
            for k,v in node.data.items():expected[k]=join(expected.get(k),v)
        diff=[i for i in range(2**depth)if x.records[i]!=y.records[i]];out=reconcile(x,y,*nodes)
        need(out['different_buckets']==diff and nodes[0].data==nodes[1].data==expected,'literal full-state reconciliation')
        need(out['compared']<=min(2**(depth+1)-1,1+2*depth*len(diff)),'difference path comparison bound')
        eq=reconcile(Snapshot(nodes[0],config),Snapshot(nodes[1],config),*nodes);need(eq['compared']==1 and not eq['different_buckets'],'equal new roots')
        counts['snapshot_pairs']+=1;counts['changed_buckets']+=len(diff)
    return counts

def encode(x):
    if isinstance(x,bytes):return x.hex()
    if hasattr(x,'__dataclass_fields__'):return asdict(x)
    raise TypeError(type(x).__name__)
if __name__=='__main__':print(json.dumps(dict(status='PASS',demonstration=main_trace(),checks=tests()),default=encode,ensure_ascii=False,indent=2))
