#!/usr/bin/env python3
"""Exact finite-permutation/Fraction models; floating CWS is a numerical teaching model.
A stored finite permutation is shared randomness, not O(k) sketch state.
"""
from fractions import Fraction as F
from itertools import permutations,product,combinations
from collections import Counter
from pathlib import Path
import math,heapq,random,json

def check(ok,message):
    if not ok:raise RuntimeError(message)

class Permutation:
    def __init__(self,order):
        order=tuple(order);N=len(order)
        if set(order)!=set(range(N)) or any(type(x) is not int for x in order):raise ValueError('permutation of 0..N-1')
        ranks=[0]*N
        for r,x in enumerate(order,1):ranks[x]=r
        self.order=order;self.ranks=tuple(ranks);self.N=N
    def rank(self,key):
        if type(key) is not int or not 0<=key<self.N:raise ValueError('key outside universe')
        return self.ranks[key]

def minhash(stream,permutation):
    best=None;rank=permutation.N+1
    for key in stream:
        r=permutation.rank(key)
        if r<rank:best=key;rank=r
    return best

class BottomK:
    def __init__(self,k,permutation):
        if type(k) is not int or k<1:raise ValueError('positive sample capacity')
        self.k=k;self.permutation=permutation;self.heap=[];self.kept=set();self.complete=True
    def add(self,key):
        rank=self.permutation.rank(key)
        if key in self.kept:return
        if len(self.heap)<self.k:heapq.heappush(self.heap,(-rank,key));self.kept.add(key);return
        self.complete=False
        if rank<-self.heap[0][0]:
            old=heapq.heapreplace(self.heap,(-rank,key))[1];self.kept.remove(old);self.kept.add(key)
    def ordered(self):return sorted(self.kept,key=self.permutation.rank)
    def cardinality(self):
        if self.complete:return F(len(self.heap))
        if self.k<2:raise ValueError('nontrivial unbiased cardinality needs k>=2')
        R=-self.heap[0][0]
        return F(self.permutation.N*(self.k-1),R-1)
    def compatible(self,other):
        # Exact reference implementation compares full shared permutations.
        if self.k!=other.k or self.permutation.ranks!=other.permutation.ranks:raise ValueError('incompatible capacity / permutation')
    def merge(self,other):
        self.compatible(other);out=BottomK(self.k,self.permutation);candidates=self.kept|other.kept
        for key in candidates:out.add(key)
        out.complete=self.complete and other.complete and len(candidates)<=self.k
        return out
    def jaccard(self,other):
        merged=self.merge(other)
        if not merged.kept:return F(1)
        return F(len(merged.kept&self.kept&other.kept),len(merged.kept))
    def verify(self,truth):
        expected=sorted(set(truth),key=self.permutation.rank)[:self.k]
        check(self.ordered()==expected,'bottom-k order');check(len(self.heap)==len(self.kept),'heap distinct')
        check(self.complete==(len(set(truth))<=self.k),'completeness bit')

def bottom(stream,k,permutation):
    out=BottomK(k,permutation)
    for key in stream:out.add(key)
    return out

class PrioritySample:
    """Input records have distinct stable IDs and immutable nonnegative weights.
    Update does not keep an O(n) seen-ID map: uniqueness is an input contract.
    """
    def __init__(self,k):
        if type(k) is not int or k<1:raise ValueError('positive k')
        self.k=k;self.heap=[]
    def add(self,key,weight,u):
        weight=F(weight);u=F(u)
        if type(key) is not int or weight<0 or not 0<u<1:raise ValueError('integer ID, nonnegative weight, 0<u<1')
        if weight==0:return
        q=weight/u;row=(q,-key,weight,u)
        if len(self.heap)<self.k+1:heapq.heappush(self.heap,row)
        elif row[:2]>self.heap[0][:2]:heapq.heapreplace(self.heap,row)
    def records(self):return sorted(self.heap,reverse=True)
    def estimates(self):
        rows=self.records();tau=rows[self.k][0] if len(rows)>self.k else F(0)
        selected={-row[1]:max(row[2],tau) for row in rows[:self.k]}
        variance={-row[1]:tau*max(F(0),tau-row[2]) for row in rows[:self.k]}
        return selected,tau,variance
    def merge_disjoint(self,other):
        if self.k!=other.k:raise ValueError('incompatible k')
        out=PrioritySample(self.k)
        # Local top-(k+1) suffices; disjoint record IDs are a precondition.
        for row in self.heap+other.heap:out.add(-row[1],row[2],row[3])
        return out

# Parameters are shared by coordinate, not redrawn for each input vector.
def cws(weights,parameters):
    if len(weights)!=len(parameters):raise ValueError('same coordinate universe')
    candidate=[]
    for i,(weight,param) in enumerate(zip(weights,parameters)):
        if not math.isfinite(weight) or weight<0:raise ValueError('finite nonnegative weights')
        r,c,beta=param
        if not (math.isfinite(r) and math.isfinite(c) and math.isfinite(beta) and r>0 and c>0 and 0<=beta<1):raise ValueError('positive Gamma values and beta in [0,1)')
        if weight==0:continue
        scaled=math.log(weight)/r+beta
        if not math.isfinite(scaled):raise ValueError('floating log-grid outside finite range')
        t=math.floor(scaled)
        log_y=r*(t-beta);log_a=math.log(c)-r*(t-beta+1)
        if not (math.isfinite(log_y) and math.isfinite(log_a)):raise ValueError('floating score outside finite range')
        candidate.append((log_a,i,t,log_y))
    if not candidate:return None,[]
    winner=min(candidate)
    return (winner[1],winner[2]),candidate

def jaccard(A,B):return F(len(A&B),len(A|B)) if A|B else F(1)
def subsets(N):return [set(i for i in range(N) if mask>>i&1) for mask in range(1<<N)]
def minhash_exact():
    N=4;sets=subsets(N);counts=Counter();cases=0
    for order in permutations(range(N)):
        p=Permutation(order);sig=[minhash(S,p) for S in sets]
        for a,A in enumerate(sets):
            for b,B in enumerate(sets):counts[a,b]+=sig[a]==sig[b];cases+=1
    for a,A in enumerate(sets):
        for b,B in enumerate(sets):check(F(counts[a,b],math.factorial(N))==jaccard(A,B),'exact MinHash collision')
    affine=Counter()
    for a in range(1,5):
        for b in range(5):affine[min(range(3),key=lambda x:(a*x+b)%5)]+=1
    check(len(set(affine.values()))>1,'affine is not minwise')
    p=Permutation((2,0,3,1));check(minhash([0,0,0,1,2],p)==minhash({0,1,2},p),'duplicates ignored')
    return {'set_pair_permutation_cases':cases,'N':N,'affine_mod5_minimum_counts':dict(affine),'affine_permutations':20}

def bottom_exact():
    merge_cases=0;jaccard_cases=0;cardinality_cases=0
    for N in range(1,7):
        sets=subsets(N);sums=Counter();perms=0
        for order in permutations(range(N)):
            perms+=1;p=Permutation(order)
            for k in range(2,N+1):
                for a,A in enumerate(sets):
                    obj=bottom(A,k,p);obj.verify(A);sums[k,a]+=obj.cardinality();cardinality_cases+=1
        for (k,a),total in sums.items():check(total/perms==len(sets[a]),'finite permutation cardinality unbiased')
    N=4;sets=subsets(N);stats={};naive=F(0)
    for order in permutations(range(N)):
        p=Permutation(order)
        for k in range(1,5):
            objects=[bottom(A,k,p) for A in sets]
            for a,A in enumerate(sets):
                for b,B in enumerate(sets):
                    x,y=objects[a],objects[b];merged=x.merge(y);merged.verify(A|B);merge_cases+=1
                    z=x.jaccard(y);key=(k,a,b);s,ss=stats.get(key,(F(0),F(0)));stats[key]=(s+z,ss+z*z);jaccard_cases+=1
        x=bottom({0,1,2},2,p);y=bottom({1,2,3},2,p);naive+=F(len(x.kept&y.kept),len(x.kept|y.kept))
    for (k,a,b),(s,ss) in stats.items():
        J=jaccard(sets[a],sets[b]);n=len(sets[a]|sets[b]);variance=F(0) if n<=k else J*(1-J)*F(n-k,k*(n-1))
        check(s/24==J and ss/24-J*J==variance,'bottom J expectation/variance')
    rng=random.Random(6110);stream_checks=0
    for trial in range(100):
        N=25;order=list(range(N));rng.shuffle(order);p=Permutation(order);obj=BottomK(5,p);truth=set()
        for i in range(200):
            x=rng.randrange(N);obj.add(x);truth.add(x);obj.verify(truth);stream_checks+=1
    return {'finite_cardinality_cases':cardinality_cases,'merge_cases':merge_cases,'Jaccard_cases':jaccard_cases,'stream_prefix_checks':stream_checks,'naive_overlap_ratio_mean':str(naive/24),'correct_overlap_J':'1/2'}

def priority_checks():
    rng=random.Random(6120);histories=0;conditional=0
    for trial in range(4000):
        n=rng.randrange(1,21);k=rng.randrange(1,8)
        records=[(i,rng.randrange(0,12),F(rng.randrange(1,101),101)) for i in range(n)]
        obj=PrioritySample(k)
        for i,(key,w,u) in enumerate(records):
            obj.add(key,w,u)
            expected=sorted([(F(ww)/uu,-kk,F(ww),uu) for kk,ww,uu in records[:i+1] if ww>0],reverse=True)[:k+1]
            check(obj.records()==expected,'priority prefix top-(k+1)')
        a,b=PrioritySample(k),PrioritySample(k)
        for i,rec in enumerate(records):(a if i%2 else b).add(*rec)
        check(a.merge_disjoint(b).records()==obj.records(),'disjoint priority merge');histories+=1
    # Integrate the one-dimensional conditional step function exactly.
    for w in range(1,10):
        for num in range(1,21):
            for den in range(1,10):
                theta=F(num,den);p=min(F(1),F(w)/theta);value=max(F(w),theta)
                check(p*value==w,'conditional unbiased integral')
                check(p*(theta*max(F(0),theta-w))==p*value*value-w*w,'conditional variance integral');conditional+=1
    light_total=F(0);heavy_total=F(0)
    # Deliberately incorrect finite random grid, used only as a counterexample.
    for u,v in product((F(1,3),F(2,3)),repeat=2):
        o=PrioritySample(1);o.add(0,1,u);o.add(1,4,v);weights,_,_=o.estimates();light_total+=weights.get(0,0);heavy_total+=weights.get(1,0)
    check(light_total==0,'finite-grid light item never enters')
    return {'stream_histories':histories,'exact_conditional_integrals':conditional,'finite_grid_counterexample':{'weights':[1,4],'u_grid':['1/3','2/3'],'mean_light_estimate':str(light_total/4),'mean_heavy_estimate':str(heavy_total/4)}}

def cws_checks():
    rng=random.Random(6130);consistency=0;union_checks=0;comparisons=0;collisions=0
    A=[1,2,0,4];B=[2,1,3,4];truth=F(sum(min(a,b) for a,b in zip(A,B)),sum(max(a,b) for a,b in zip(A,B)))
    for trial in range(20000):
        params=[(rng.gammavariate(2,1),rng.gammavariate(2,1),rng.random()) for _ in range(4)]
        S=[rng.randrange(0,10) for _ in range(4)];T=[rng.randrange(x+1) for x in S]
        sig,rows=cws(S,params);small,_=cws(T,params)
        if sig is not None:
            i=sig[0];ly=next(row[3] for row in rows if row[1]==i)
            if T[i]>0 and ly<math.log(T[i])-1e-11:check(sig==small,'CWS shrink consistency');consistency+=1
        U=[max(a,b) for a,b in zip(A,B)];su,ru=cws(U,params);sa,_=cws(A,params);sb,_=cws(B,params)
        row=next(row for row in ru if row[1]==su[0]);bound=min(A[su[0]],B[su[0]])
        inside=bound>0 and row[3]<=math.log(bound)
        check((sa==sb)==inside,'CWS union/intersection equivalence');union_checks+=1
        comparisons+=1;collisions+=sa==sb
    # Exact factorized mixed moments of V=RB and W=R(1-B), Gamma/Beta integrals.
    moments=0
    for p in range(7):
        for q in range(7):
            gamma_moment=math.factorial(p+q+1);beta_integral=F(math.factorial(p)*math.factorial(q),math.factorial(p+q+1))
            check(gamma_moment*beta_integral==math.factorial(p)*math.factorial(q),'V/W product moments');moments+=1
    params=[(math.log(4),2,.5),(math.log(4),1,.5)]
    signatures=[cws(v,params)[0] for v in ([1,5],[1,3],[1,1])]
    check(signatures==[(1,1),(1,1),(1,0)],'CWS structural hand calculation')
    check(cws([0,0],params)[0] is None,'empty CWS')
    return {'numerical_shrink_consistency_checks':consistency,'numerical_union_equivalence_checks':union_checks,'exact_Gamma_Beta_moment_checks':moments,'hand_signatures':signatures,'numerical_collision_diagnostic':{'trials':comparisons,'matches':collisions,'frequency':collisions/comparisons,'ideal_weighted_J':str(truth),'note':'Numerical diagnostic only; floating PRNG samples do not prove exact continuous probabilities.'}}

def examples():
    p=Permutation((0,1,2,3,4,5));A={0,1,2};B={1,2,3};a=bottom(A,2,p);b=bottom(B,2,p)
    pr=PrioritySample(2)
    records=[(0,2,F(1,2)),(1,5,F(1,2)),(2,1,F(1,5)),(3,3,F(3,4))]
    for row in records:pr.add(*row)
    est,tau,var=pr.estimates()
    return {'universe_order':p.order,'A':sorted(A),'B':sorted(B),'minhash_A':minhash(A,p),'minhash_B':minhash(B,p),'bottom_A':a.ordered(),'bottom_B':b.ordered(),'bottom_union':a.merge(b).ordered(),'bottom_J':str(a.jaccard(b)),'cardinality_A':str(a.cardinality()),'cardinality_B':str(b.cardinality()),'priority_records':[[i,w,str(u)] for i,w,u in records],'priority_threshold':str(tau),'priority_estimates':{str(k):str(v) for k,v in est.items()},'priority_variance_estimates':{str(k):str(v) for k,v in var.items()}}

def main():
    result={'result':'PASS','examples':examples(),'minhash':minhash_exact(),'bottom_k':bottom_exact(),'priority':priority_checks(),'CWS':cws_checks()}
    Path(__file__).with_name('algorithms-coordinated-sampling-results.json').write_text(json.dumps(result,ensure_ascii=False,indent=2)+'\n')
    print(json.dumps(result,ensure_ascii=False,indent=2))
if __name__=='__main__':main()
