#!/usr/bin/env python3
"""Finite poset, nerve and Dowker certificates, using only the standard library.
Integer chain identities and GF(2) homology are separate checks. Contractibility
is certified by actual comparable maps, not inferred from vanishing Betti numbers.
"""
from itertools import combinations, product
from collections import Counter
from pathlib import Path
import argparse, json, copy

COUNTS=Counter()
def require(ok,label):
    if not ok: raise ValueError(label)
def check(ok,label):
    require(ok,label);COUNTS[label]+=1

def subsets(xs):
    xs=tuple(xs)
    for r in range(1,len(xs)+1):
        yield from combinations(xs,r)
def complex_of(facets):
    return {s for facet in facets for s in subsets(sorted(set(facet)))}
def mask(s):return sum(1<<x for x in s)
def bits(m):return tuple(i for i in range(m.bit_length()) if m>>i&1)
def face_masks(K):return {mask(s) for s in K}
def is_complex(K):
    return all(tuple(sorted(set(s)))==s and s and all(t in K for t in subsets(s)) for s in K)

class Poset:
    def __init__(self,nodes,le):
        self.nodes=tuple(sorted(nodes));self.le={(x,y) for x in self.nodes for y in self.nodes if le(x,y)}
        require(all((x,x)in self.le for x in self.nodes),'Poset reflexivity')
        require(all(x==y or (y,x)not in self.le for x,y in self.le),'Poset antisymmetry')
        require(all((x,z)in self.le for x,y in self.le for y2,z in self.le if y==y2),'Poset transitivity')
        self._complex=None
    def below(self,x,y):return (x,y)in self.le
    def sub(self,nodes):return Poset(nodes,self.below)
    def opposite(self):return Poset(self.nodes,lambda x,y:self.below(y,x))
    def chains(self):
        if self._complex is not None:return self._complex
        out=set()
        def grow(path):
            out.add(tuple(sorted(path)))
            for y in self.nodes:
                if path[-1]!=y and self.below(path[-1],y):grow(path+(y,))
        for x in self.nodes:grow((x,))
        self._complex=out;return out
    def linear_extension(self):
        todo=set(self.nodes);out=[]
        while todo:
            x=min(x for x in todo if not any(y!=x and self.below(y,x)for y in todo))
            out.append(x);todo.remove(x)
        return out

def faces_poset(K):
    require(is_complex(K),'Face closure')
    return Poset(face_masks(K),lambda x,y:x&y==x)
def monotone(P,Q,f):
    return set(f)==set(P.nodes) and all(f[x]in Q.nodes for x in P.nodes) and all(Q.below(f[x],f[y])for x,y in P.le)
def oriented(vertices):
    if len(set(vertices))!=len(vertices):return None,0
    inversions=sum(vertices[i]>vertices[j] for i in range(len(vertices))for j in range(i+1,len(vertices)))
    return tuple(sorted(vertices)),(-1)**inversions

def add(chain,s,c):
    if c and s is not None:
        chain[s]=chain.get(s,0)+c
        if not chain[s]:del chain[s]
def boundary(chain):
    out={}
    for s,c in chain.items():
        if len(s)>1:
            for i in range(len(s)):add(out,s[:i]+s[i+1:],c*(-1)**i)
    return out

def image(chain,f):
    out={}
    for s,c in chain.items():
        t,sign=oriented(tuple(f[x]for x in s));add(out,t,c*sign)
    return out

def difference(a,b):
    out=dict(a)
    for s,c in b.items():add(out,s,-c)
    return out

def chain_map_check(K,L,f):
    require(set(f)=={x for s in K for x in s},'Vertex-map domain')
    for s in K:
        t,sign=oriented(tuple(f[x]for x in s))
        require(t is None or t in L,'Invalid simplex image')
        check(boundary(image({s:1},f))==image(boundary({s:1}),f),'Integer chain-map identity')

def prism(s,f,g,L):
    out={}
    for j in range(len(s)):
        vertices=tuple(f[x]for x in s[:j+1])+tuple(g[x]for x in s[j:])
        t,sign=oriented(vertices)
        require(t is None or t in L,'Prism target simplex')
        add(out,t,sign*(-1)**j)
    return out

def order_homotopy(P,Q,f,g,check_chains=True):
    require(monotone(P,Q,f)and monotone(P,Q,g),'Order homotopy maps must be monotone')
    require(all(Q.below(f[x],g[x])for x in P.nodes),'Wrong pointwise comparison direction')
    current=dict(f);steps=[];K=P.chains();L=Q.chains();total={s:{}for s in K}
    while current!=g:
        changed={x for x in P.nodes if current[x]!=g[x]}
        x=min(x for x in changed if not any(y!=x and P.below(x,y)for y in changed))
        nxt=dict(current);nxt[x]=g[x]
        check(monotone(P,Q,nxt),'Single-vertex update remains monotone')
        if check_chains:
            for s in K:
                union=tuple(sorted({current[v]for v in s}|{nxt[v]for v in s}))
                check(union in L,'Each update is simplicially contiguous')
                for t,c in prism(s,current,nxt,L).items():add(total[s],t,c)
        steps.append({'vertex':x,'before':current[x],'after':nxt[x]});current=nxt
    if check_chains:
        for s in K:
            lhs=boundary(total[s])
            for face,c in boundary({s:1}).items():
                for t,k in total[face].items():add(lhs,t,c*k)
            rhs=difference(image({s:1},g),image({s:1},f))
            check(lhs==rhs,'Integer prism chain-homotopy identity')
    return {'steps':steps,'changed_vertices':len(steps),
            'chain_homotopy':[{ 'source':list(s),'terms':[{'simplex':list(t),'coefficient':c}for t,c in sorted(v.items())]}
                               for s,v in sorted(total.items(),key=lambda z:(len(z[0]),z[0]))]if check_chains else []}

def contraction_check(P,maps):
    require(P.nodes,'Empty realization is not contractible')
    require(len(maps)>=1 and maps[0]=={x:x for x in P.nodes},'Contraction must start at identity')
    require(len(set(maps[-1].values()))==1,'Contraction must end at a point')
    for h in maps:require(monotone(P,P,h),'Contraction map invalid')
    for f,g in zip(maps,maps[1:]):
        require(all(P.below(f[x],g[x])for x in P.nodes)or all(P.below(g[x],f[x])for x in P.nodes),'Contraction maps incomparable')
        check(True,'Contractibility certified by pointwise-comparable maps')
    return True

def cone_face_contraction(K):
    require(K,'Empty complex cannot supply cone vertex')
    vertices=sorted({v for s in K for v in s})
    v=next((v for v in vertices if all(tuple(sorted(set(s)|{v}))in K for s in K)),None)
    require(v is not None,'No cone-vertex certificate for this intersection')
    P=faces_poset(K);maps=[{x:x for x in P.nodes},{x:x|(1<<v)for x in P.nodes},{x:1<<v for x in P.nodes}]
    contraction_check(P,maps)
    return maps, v

def quillen(P,Q,f,contractions):
    require(monotone(P,Q,f),'Quillen input not monotone')
    require(set(contractions)==set(Q.nodes),'One contraction per target required')
    fibers={}
    for y in Q.nodes:
        fiber={x for x in P.nodes if Q.below(f[x],y)}
        contraction_check(P.sub(fiber),contractions[y]);fibers[y]=fiber
    pindex={x:i for i,x in enumerate(P.nodes)};qindex={y:len(P.nodes)+i for i,y in enumerate(Q.nodes)}
    backp={v:k for k,v in pindex.items()};backq={v:k for k,v in qindex.items()}
    def le(x,y):
        if x in backp and y in backp:return P.below(backp[x],backp[y])
        if x in backq and y in backq:return Q.below(backq[x],backq[y])
        if x in backp and y in backq:return Q.below(f[backp[x]],backq[y])
        return False
    M=Poset(list(backp)+list(backq),le);logs=[];current=set(M.nodes)
    for x in reversed(P.linear_extension()):
        node=pindex[x];upper={y for y in current if y!=node and M.below(node,y)};pivot=qindex[f[x]]
        check(pivot in upper and all(M.below(pivot,y)for y in upper),'Mapping cylinder source deletion upper minimum')
        logs.append({'side':'source','original':x,'witness':f[x],'current_size':len(current)})
        current.remove(node)
    check(current==set(backq),'Source deletion leaves target')
    current=set(M.nodes)
    for y in Q.linear_extension():
        node=qindex[y];lower={x for x in current if x!=node and M.below(x,node)}
        check(lower=={pindex[x]for x in fibers[y]},'Mapping cylinder target deletion exact lower fiber')
        logs.append({'side':'target','original':y,'fiber':sorted(fibers[y]),'current_size':len(current)})
        current.remove(node)
    check(current==set(backp),'Target deletion leaves source')
    return logs

def boundary_columns(K,d):
    high=sorted(s for s in K if len(s)==d+1);low=sorted(s for s in K if len(s)==d)
    lookup={s:i for i,s in enumerate(low)}
    cols=[]
    for s in high:
        c=0
        if d:
            for i in range(len(s)):c^=1<<lookup[s[:i]+s[i+1:]]
        cols.append(c)
    return high,cols

def rank2(columns):
    piv={}
    for col in columns:
        while col:
            i=col.bit_length()-1
            if i in piv:col^=piv[i]
            else:piv[i]=col;break
    return len(piv)

def cycles2(columns):
    piv={};basis=[]
    for j,col in enumerate(columns):
        track=1<<j
        while col:
            i=col.bit_length()-1
            if i in piv:col^=piv[i][0];track^=piv[i][1]
            else:piv[i]=(col,track);break
        if not col:basis.append(track)
    return basis

def betti(K):
    maxd=max([len(s)-1 for s in K]+[0]);out=[]
    for d in range(maxd+1):
        high,cols=boundary_columns(K,d);_,nextcols=boundary_columns(K,d+1)
        out.append(len(high)-rank2(cols)-rank2(nextcols))
    while len(out)>1 and out[-1]==0:out.pop()
    return out

def map_ranks(K,L,f):
    maxd=max([len(s)-1 for s in K|L]+[0]);out=[]
    for d in range(maxd+1):
        sources,cols=boundary_columns(K,d);targets,_=boundary_columns(L,d);lookup={s:i for i,s in enumerate(targets)}
        image_cols=[]
        for s in sources:
            t,sign=oriented(tuple(f[x]for x in s))
            require(t is None or t in lookup,'Homology map invalid simplex')
            image_cols.append(0 if t is None else 1<<lookup[t])
        images=[]
        for cycle in cycles2(cols):
            result=0
            for i in bits(cycle):result^=image_cols[i]
            images.append(result)
        _,boundaries=boundary_columns(L,d+1)
        out.append(rank2(boundaries+images)-rank2(boundaries))
    while len(out)>1 and out[-1]==0:out.pop()
    return out

def common_columns(R,s,n):return mask(y for y in range(n)if all((x,y)in R for x in bits(s)))
def common_rows(R,t,m):return mask(x for x in range(m)if all((x,y)in R for y in bits(t)))
def dowker(m,n,R):
    require(all(0<=x<m and 0<=y<n for x,y in R),'Relation label outside declared sets')
    X={bits(s)for s in range(1,1<<m)if common_columns(R,s,n)}
    Y={bits(t)for t in range(1,1<<n)if common_rows(R,t,m)}
    P=faces_poset(X);Q=faces_poset(Y).opposite();f={s:common_columns(R,s,n)for s in P.nodes}
    require(monotone(P,Q,f),'Common-neighbor order direction')
    contractions={};maxima=[]
    for t in Q.nodes:
        pivot=common_rows(R,t,m);fiber={s for s in P.nodes if Q.below(f[s],t)}
        expected={s for s in range(1,1<<m)if s&pivot==s}
        check(fiber==expected and pivot in fiber,'Dowker exact fiber and maximum')
        contractions[t]=[{s:s for s in fiber},{s:pivot for s in fiber}]
        maxima.append({'target_face':list(bits(t)),'maximum_source_face':list(bits(pivot))})
    logs=quillen(P,Q,f,contractions) if R else []
    chain_map_check(P.chains(),Q.chains(),f)
    ranks=map_ranks(P.chains(),Q.chains(),f)
    check(ranks==betti(P.chains())==betti(Q.chains()),'Canonical Dowker map induces full homology isomorphism')
    check(betti(X)==betti(Y),'Original Dowker Betti numbers agree')
    return {'X':X,'Y':Y,'P':P,'Q':Q,'map':f,'fiber_maxima':maxima,'cylinder_log':logs,'homology_ranks':ranks}

def beat_deletions(P,steps):
    current=P;out=[]
    for x,direction,witness in steps:
        require(x in current.nodes and witness in current.nodes and x!=witness,'Deleted or absent beat witness')
        require(direction in ('down','up'),'Beat direction')
        neighbors=[y for y in current.nodes if y!=x and
                   (current.below(y,x)if direction=='down'else current.below(x,y))]
        require(witness in neighbors and all(current.below(y,witness)if direction=='down'else current.below(witness,y)
                                             for y in neighbors),'Beat extremum missing')
        target=current.sub([y for y in current.nodes if y!=x])
        r={y:witness if y==x else y for y in current.nodes};identity={y:y for y in current.nodes}
        check(monotone(current,target,r),'Beat retraction is monotone')
        certificate=order_homotopy(current,current,r,identity)if direction=='down'else order_homotopy(current,current,identity,r)
        check(all(r[y]==y for y in target.nodes),'Beat homotopy fixes remaining vertices')
        out.append({'current_vertices':list(current.nodes),'deleted':x,'direction':direction,'witness':witness,
                    'strict_neighbors':neighbors,'retraction':[[y,r[y]]for y in current.nodes],
                    'integer_chain_homotopy':certificate})
        current=target
    return out

def nerve_cover(K,members):
    require(is_complex(K),'Original complex invalid')
    require(all(is_complex(L)and L<=K for L in members),'Cover member is not a subcomplex')
    require(set().union(*members)==K,'Cover misses original simplices')
    intersections={};N=set();contractions={}
    for J in range(1,1<<len(members)):
        meet=set.intersection(*(members[j]for j in bits(J)))
        if meet:
            maps,v=cone_face_contraction(meet);intersections[J]=(meet,v);contractions[J]=maps;N.add(bits(J))
    P=faces_poset(K);Q=faces_poset(N).opposite()
    f={s:mask(i for i,L in enumerate(members)if bits(s)in L)for s in P.nodes}
    logs=quillen(P,Q,f,contractions)
    chain_map_check(P.chains(),Q.chains(),f)
    ranks=map_ranks(P.chains(),Q.chains(),f)
    check(ranks==betti(K)==betti(N),'Nerve canonical map homology isomorphism')
    return {'nerve':N,'map':f,'intersection_cones':[{'labels':list(bits(j)),'cone_vertex':v,'nonempty_faces':len(KJ)}for j,(KJ,v)in sorted(intersections.items())],
            'cylinder_log':logs,'homology_ranks':ranks}

def rejection(thunk,label):
    try:thunk()
    except ValueError:COUNTS[label]+=1;return
    raise ValueError('Broken input accepted: '+label)

def summary(d):
    def facets(K):return [list(s)for s in sorted(K)if not any(set(s)<set(t)for t in K)]
    return {'row_facets':facets(d['X']),'column_facets':facets(d['Y']),
       'row_f_vector':[sum(len(s)==k for s in d['X'])for k in range(1,1+max([len(s)for s in d['X']]+[0]))],
       'column_f_vector':[sum(len(s)==k for s in d['Y'])for k in range(1,1+max([len(s)for s in d['Y']]+[0]))],
       'row_betti':betti(d['X']),'column_betti':betti(d['Y']),
       'face_map':[{'source':list(bits(s)),'target':list(bits(t))}for s,t in sorted(d['map'].items())],
       'fiber_maxima':d['fiber_maxima'],'map_homology_ranks':d['homology_ranks'],'mapping_cylinder_deletions':d['cylinder_log']}

def run():
    COUNTS.clear()
    P=Poset([0,1],lambda x,y:x<=y)
    Q=Poset([0,1,2,3],lambda x,y:x==y or x==0 or y==3)
    first=order_homotopy(P,Q,{0:0,1:1},{0:2,1:3})
    check(first['changed_vertices']==2,'Two-stage diamond homotopy')
    beatlog=beat_deletions(Q,[(1,'down',0),(2,'down',0),(3,'down',0)])
    rejection(lambda:beat_deletions(Q,[(1,'down',0),(2,'down',1)]),'Deleted beat witness rejected')
    rejection(lambda:beat_deletions(Poset([0,1],lambda x,y:x==y),[(0,'up',1)]),'Isolated beat point rejected')
    finite_posets=0;beat_cases=0
    for n in range(1,5):
        pairs=list(combinations(range(n),2))
        for choices in product(range(3),repeat=len(pairs)):
            rel={(x,x)for x in range(n)}
            for (a,b),v in zip(pairs,choices):
                if v:rel.add((a,b)if v==1 else (b,a))
            if not all((x,z)in rel for x,y in rel for y2,z in rel if y==y2):continue
            T=Poset(range(n),lambda x,y:(x,y)in rel);finite_posets+=1
            for x in T.nodes:
                for direction in ('down','up'):
                    near=[y for y in T.nodes if y!=x and (T.below(y,x)if direction=='down'else T.below(x,y))]
                    for w in near:
                        if all(T.below(y,w)if direction=='down'else T.below(w,y)for y in near):
                            beat_deletions(T,[(x,direction,w)]);beat_cases+=1
    circle=complex_of([(0,1),(1,2),(0,2)])
    good=nerve_cover(circle,[complex_of([e])for e in [(0,1),(1,2),(0,2)]])
    rejection(lambda:nerve_cover(circle,[complex_of([(0,1),(1,2)]),complex_of([(0,2)])]),'Disconnected pairwise intersection rejected')
    rejection(lambda:nerve_cover(complex_of([(0,1)]),[complex_of([(0,)]),complex_of([(1,)])]),'Vertex-only cover rejected')
    A=Poset([0,1],lambda x,y:x==y);B=Poset([0,1],lambda x,y:x<=y)
    wrong={0:[{0:0},{0:0}],1:[{1:1},{1:1}]}
    rejection(lambda:quillen(A,B,{0:0,1:1},wrong),'Exact-point fibers cannot replace lower fibers')
    rejection(lambda:contraction_check(Poset([],lambda x,y:False),[{}]),'Empty contraction rejected')
    R0={(0,0),(0,1),(1,1),(1,2),(2,0),(2,2),(3,0),(3,1)}
    relations=[R0,R0|{(0,2)},R0|{(0,2),(3,2)}]
    models=[dowker(4,3,R)for R in relations];naturality=[]
    for R,S,d,e in zip(relations,relations[1:],models,models[1:]):
        f={s:common_columns(S,s,3)for s in d['P'].nodes};g=d['map']
        certificate=order_homotopy(d['P'],e['Q'],f,g)
        incX={v:v for s in d['X']for v in s};incY={v:v for s in d['Y']for v in s}
        chain_map_check(d['X'],e['X'],incX);chain_map_check(d['Y'],e['Y'],incY)
        rx=map_ranks(d['X'],e['X'],incX);ry=map_ranks(d['Y'],e['Y'],incY)
        check(rx==ry,'Both relation-filtration inclusion ranks agree')
        naturality.append({'row_inclusion_ranks':rx,'column_inclusion_ranks':ry,'homotopy':certificate})
    # All 512 relations, including isolated labels and the empty relation.
    hist=Counter()
    for code in range(1<<9):
        R={(i,j)for i,j in product(range(3),repeat=2)if code>>(3*i+j)&1}
        d=dowker(3,3,R);hist[','.join(map(str,betti(d['X'])))]+=1
    # Duplicate a witness column; the row complex must remain exact, not just homologous.
    for code in range(1<<6):
        R={(i,j)for i,j in product(range(3),range(2))if code>>(2*i+j)&1}
        S=R|{(i,2)for i in range(3)if (i,0)in R}
        d=dowker(3,2,R);e=dowker(3,3,S)
        check(d['X']==e['X'],'Duplicating a column preserves exact row faces')
    # Witness data must come from the original relation, including high-dimensional faces.
    row_faces=models[0]['X'];pairwise_clique=(0,1,2)
    check(all(tuple(pair)in row_faces for pair in combinations(pairwise_clique,2))and pairwise_clique not in row_faces,
          'Pairwise witnesses do not create a triple face')
    corrupted=copy.deepcopy(models[0]['map']);corrupted[mask([0])]=mask([0,1,2])
    rejection(lambda: require(corrupted=={s:common_columns(R0,s,3)for s in models[0]['P'].nodes},'False common-neighbor table'),
              'Altered common-neighbor map rejected')
    corrupted=copy.deepcopy(models[0]['fiber_maxima']);corrupted[0]['maximum_source_face']=[1]
    rejection(lambda:require(all(mask(v['maximum_source_face'])==common_rows(R0,mask(v['target_face']),4)for v in corrupted),'False fiber maximum'),
              'Altered fiber maximum rejected')
    return {'status':'PASS','checks':sum(COUNTS.values()),'categories':dict(COUNTS),
            'coefficient_contract':'Integer chain-map and chain-homotopy identities; Betti numbers and induced ranks over F2 only.',
            'diamond_homotopy':first,'diamond_beat_deletions':beatlog,'all_labeled_posets_sizes_one_to_four':finite_posets,'individual_beat_certificates':beat_cases,'good_circle_cover':{
                'nerve_faces':[list(s)for s in sorted(good['nerve'])],
                'intersection_cones':good['intersection_cones'],'map_homology_ranks':good['homology_ranks']},
            'relation_filtration':[{'relation':sorted(map(list,R)),**summary(d)}for R,d in zip(relations,models)],
            'naturality_certificates':naturality,'all_3_by_3_relations':512,'betti_histogram':dict(sorted(hist.items())),
            'duplicate_column_models':64,
            'scope':'Finite constructive certificates. No inference from vanishing Betti numbers to contractibility; no general contractibility algorithm, arbitrary open-cover theorem, or claim about integer torsion from F2 ranks.'}

if __name__=='__main__':
    p=argparse.ArgumentParser(description=__doc__);p.add_argument('--output',type=Path);args=p.parse_args()
    result=run();encoded=json.dumps(result,ensure_ascii=False,indent=2)+'\n'
    if args.output:args.output.write_text(encoded)
    print(json.dumps({'status':result['status'],'checks':result['checks']}))
