#!/usr/bin/env python3
"""Small exact enumeration implementations and explicit self-checks (also under -O)."""
from bisect import insort
from collections import deque
from dataclasses import dataclass
from heapq import heappush,heappop,heapify
from itertools import combinations,product
from math import inf
from pathlib import Path
import json,random

def check(ok,label):
    if not ok:raise AssertionError(label)

def lawler(bits,opt,limit):
    """opt(prefix)->(cost, full bit tuple) or None; oracle must be an exact minimizer."""
    if type(bits) is not int or bits<0 or type(limit) is not int or limit<0:raise ValueError('bad dimension/limit')
    if not limit:return [],{'oracle_calls':0,'first_partition':[]}
    first=opt(());calls=1
    if first is None:return [],{'oracle_calls':calls,'first_partition':[]}
    heap=[(first[0],0,(),first[1])];serial=1;out=[];partition=[]
    while heap and len(out)<limit:
        cost,_,prefix,x=heappop(heap);check(len(x)==bits and x[:len(prefix)]==prefix,'oracle witness');out.append((cost,x))
        if len(out)==limit:break
        for j in range(len(prefix),bits):
            child=x[:j]+(1-x[j],);best=opt(child);calls+=1
            if len(out)==1:partition.append((child,best))
            if best is not None:
                heappush(heap,(best[0],serial,child,best[1]));serial+=1
    return out,{'oracle_calls':calls,'first_partition':partition}

def validate_graph(n,edges):
    if type(n) is not int or n<1:raise ValueError('nonempty graph required')
    for u,v,w in edges:
        if type(u) is not int or type(v) is not int or not 0<=u<n or not 0<=v<n or type(w) is not int or w<0:raise ValueError('bad nonnegative integer arc')

def shortest_spur(n,edges,s,t,ban_vertices=frozenset(),ban_edges=frozenset()):
    if s in ban_vertices or t in ban_vertices:return None
    adj=[[] for _ in range(n)]
    for e,(u,v,w) in enumerate(edges):
        if e not in ban_edges and u not in ban_vertices and v not in ban_vertices:adj[u].append((v,w,e))
    distance=[inf]*n;parent=[None]*n;distance[s]=0;heap=[(0,s)]
    while heap:
        d,u=heappop(heap)
        if d!=distance[u]:continue
        if u==t:
            path=[];vertices=[t]
            while u!=s:
                u,e=parent[u];path.append(e);vertices.append(u)
            return d,tuple(reversed(path)),tuple(reversed(vertices))
        for v,w,e in adj[u]:
            if d+w<distance[v]:distance[v]=d+w;parent[v]=(u,e);heappush(heap,(d+w,v))
    return None

def yen(n,edges,s,t,limit):
    validate_graph(n,edges)
    if type(limit) is not int or limit<0 or type(s) is not int or type(t) is not int or not 0<=s<n or not 0<=t<n:raise ValueError('bad endpoints/limit')
    if limit==0:return [],{'spur_calls':0}
    first=shortest_spur(n,edges,s,t);calls=1
    if first is None:return [],{'spur_calls':calls}
    accepted=[first];heap=[];known={first[1]};blocked_next={};serial=0
    def register(path):
        for j,e in enumerate(path):blocked_next.setdefault(path[:j],set()).add(e)
    register(first[1])
    while len(accepted)<limit:
        _,path,vertices=accepted[-1];root_cost=0
        for j,e in enumerate(path):
            prefix=path[:j];spur=shortest_spur(n,edges,vertices[j],t,set(vertices[:j]),blocked_next[prefix]);calls+=1
            if spur is not None:
                cost,tail,tail_vertices=spur;candidate=prefix+tail
                if candidate not in known:
                    known.add(candidate);heappush(heap,(root_cost+cost,serial,candidate,vertices[:j]+tail_vertices));serial+=1
            root_cost+=edges[e][2]
        if not heap:break
        cost,_,path,vertices=heappop(heap);accepted.append((cost,path,vertices));register(path)
    return accepted,{'spur_calls':calls,'generated_candidates':len(known)}

@dataclass(frozen=True,slots=True)
class HeapNode:
    item:tuple
    left:object=None
    right:object=None

def persistent_insert(root,size,item):
    """Path-copy a complete binary heap; old roots remain unchanged."""
    direction=bin(size+1)[3:]
    def put(node,depth,carry):
        if depth==len(direction):
            check(node is None,'complete heap insertion slot');return HeapNode(carry)
        check(node is not None,'complete heap ancestor')
        keep,carry=(carry,node.item) if carry<node.item else (node.item,carry)
        if direction[depth]=='0':return HeapNode(keep,put(node.left,depth+1,carry),node.right)
        return HeapNode(keep,node.left,put(node.right,depth+1,carry))
    return put(root,0,item)

def eppstein_index(n,edges,s,t,limit):
    """Return implicit records (cost,prefix_record,last_sidetrack), plus decoding data."""
    validate_graph(n,edges)
    if type(limit) is not int or limit<0 or type(s) is not int or type(t) is not int or not 0<=s<n or not 0<=t<n:raise ValueError('bad endpoints/limit')
    if not limit:return [],{'tree':[None]*n,'distances':[inf]*n,'heap_nodes':0}
    rev=[[] for _ in range(n)];outgoing=[[] for _ in range(n)]
    for e,(u,v,w) in enumerate(edges):rev[v].append((u,w,e));outgoing[u].append(e)
    d=[inf]*n;tree=[None]*n;d[t]=0;pq=[(0,t)];order=[]
    while pq:
        value,u=heappop(pq)
        if value!=d[u]:continue
        order.append(u)
        for v,w,e in rev[u]:
            if value+w<d[v]:d[v]=value+w;tree[v]=e;heappush(pq,(d[v],v))
    if d[s]==inf:return [],{'tree':tree,'distances':d,'heap_nodes':0}
    local=[[] for _ in range(n)];roots=[None]*n;sizes=[0]*n;delta={};new_nodes=0
    for u in order:
        items=[]
        for e in outgoing[u]:
            _,v,w=edges[e]
            if d[v]!=inf:
                delta[e]=w+d[v]-d[u];check(delta[e]>=0,'nonnegative sidetrack difference')
                if e!=tree[u]:items.append((delta[e],e))
        inherited=roots[edges[tree[u]][1]] if tree[u] is not None else None
        count=sizes[edges[tree[u]][1]] if tree[u] is not None else 0
        if items:
            least=min(items);local[u]=[x for x in items if x!=least];heapify(local[u])
            roots[u]=persistent_insert(inherited,count,least);sizes[u]=count+1;new_nodes+=(count+1).bit_length()
        else:roots[u]=inherited;sizes[u]=count
    def key(ref):
        return ref[1].item if ref[0]=='T' else local[ref[1]][ref[2]]
    def children(ref):
        if ref[0]=='T':
            node=ref[1]
            if node.left is not None:yield ('T',node.left)
            if node.right is not None:yield ('T',node.right)
            tail=edges[node.item[1]][0]
            if local[tail]:yield ('A',tail,0)
        else:
            _,tail,i=ref
            for j in (2*i+1,2*i+2):
                if j<len(local[tail]):yield ('A',tail,j)
    records=[(d[s],None,None)];frontier=[];serial=0
    if roots[s] is not None:
        ref=('T',roots[s]);heappush(frontier,(d[s]+key(ref)[0],serial,ref,0));serial+=1
    while frontier and len(records)<limit:
        cost,_,ref,prefix=heappop(frontier);extra,e=key(ref);index=len(records);records.append((cost,prefix,e))
        for child in children(ref):
            inc=key(child)[0]-extra;check(inc>=0,'heap replacement cost')
            heappush(frontier,(cost+inc,serial,child,prefix));serial+=1
        head=edges[e][1]
        if roots[head] is not None:
            child=('T',roots[head]);heappush(frontier,(cost+key(child)[0],serial,child,index));serial+=1
    return records,{'tree':tree,'distances':d,'delta':delta,'heap_nodes':new_nodes,'frontier_size':len(frontier),'local_heap_items':sum(map(len,local))}

def decode_walk(records,index,tree,edges,s,t):
    sidetracks=[]
    while index:
        _,index,e=records[index];sidetracks.append(e)
    sidetracks.reverse();path=[];u=s
    for e in sidetracks:
        tail,head,w=edges[e];steps=0
        while u!=tail:
            a=tree[u];check(a is not None,'sidetrack tail not on tree path');path.append(a);u=edges[a][1];steps+=1;check(steps<=len(tree),'tree cycle')
        path.append(e);u=head
    steps=0
    while u!=t:
        e=tree[u];check(e is not None,'no final tree suffix');path.append(e);u=edges[e][1];steps+=1;check(steps<=len(tree),'tree cycle')
    return tuple(path)

def tree_path(n,edges,tree,s,t):
    adj=[[] for _ in range(n)]
    for e in tree:
        u,v=edges[e];adj[u].append((v,e));adj[v].append((u,e))
    pred=[None]*n;pred[s]=(-1,None);queue=deque([s])
    while queue:
        u=queue.popleft()
        if u==t:break
        for v,e in adj[u]:
            if pred[v] is None:pred[v]=(u,e);queue.append(v)
    if pred[t] is None:return None
    path=[]
    while t!=s:t,e=pred[t];path.append(e)
    return path

def replace_edge(tree,remove,add):
    a=[e for e in tree if e!=remove];insort(a,add);return tuple(a)

def root_spanning_tree(n,edges):
    if not n:return None
    adj=[[] for _ in range(n)]
    for e,(u,v) in enumerate(edges):adj[u].append((v,e));adj[v].append((u,e))
    seen=[False]*n;seen[0]=True;count=1;chosen=[False]*len(edges);stack=[(0,iter(adj[0]))]
    while stack:
        try:v,e=next(stack[-1][1])
        except StopIteration:stack.pop();continue
        if not seen[v]:seen[v]=True;count+=1;chosen[e]=True;stack.append((v,iter(adj[v])))
    return tuple(e for e,used in enumerate(chosen) if used) if count==n else None

def first_difference(a,b):
    j=0
    for e in a:
        while j<len(b) and b[j]<e:j+=1
        if j==len(b) or b[j]!=e:return e
    return None

def reverse_search_trees(n,edges,stats=None):
    """Streaming stackless reverse search; no stored list of prior solutions."""
    if type(n) is not int or n<0:raise ValueError('bad vertex count')
    if any(type(u) is not int or type(v) is not int or not 0<=u<n or not 0<=v<n or u==v for u,v in edges) or len({tuple(sorted(e)) for e in edges})!=len(edges):raise ValueError('simple undirected graph required')
    root=root_spanning_tree(n,edges)
    if root is None:return
    rootmask=[False]*len(edges)
    for e in root:rootmask[e]=True
    width=n-1;slots=len(edges)*width
    def parent(tree):
        if tree==root:return None
        add=first_difference(root,tree);u,v=edges[add];cycle=tree_path(n,edges,tree,u,v)
        remove=max(e for e in cycle if not rootmask[e])
        result=replace_edge(tree,remove,add);check(sum(rootmask[e] for e in result)==sum(rootmask[e] for e in tree)+1,'reverse-search potential')
        return result
    def neighbor(tree,j):
        add=j//width;remove=tree[j%width]
        if add in tree:return None
        u,v=edges[add];cycle=tree_path(n,edges,tree,u,v)
        if remove not in cycle:return None
        return replace_edge(tree,remove,add)
    current=root;j=0;yield current
    while True:
        if j<slots:
            candidate=neighbor(current,j);j+=1
            if stats is not None:stats['neighbor_slots']=stats.get('neighbor_slots',0)+1
            if candidate is not None and candidate!=root and parent(candidate)==current:
                current=candidate;j=0;yield current
        else:
            if current==root:return
            child=current;current=parent(child)
            add=first_difference(child,current);remove=first_difference(current,child)
            j=add*width+current.index(remove)+1
            if stats is not None:stats['backtracks']=stats.get('backtracks',0)+1

def brute_simple(n,edges,s,t):
    adj=[[] for _ in range(n)]
    for e,(u,v,w) in enumerate(edges):adj[u].append((e,v,w))
    out=[]
    def dfs(u,seen,path,cost):
        if u==t:out.append((cost,tuple(path)));return
        for e,v,w in adj[u]:
            if v not in seen:dfs(v,seen|{v},path+[e],cost+w)
    dfs(s,{s},[],0);return sorted(out)

def brute_walk_costs(n,edges,s,t,k):
    # Independent k-pop Dijkstra on original walk prefixes; includes continuations after t.
    adj=[[] for _ in range(n)]
    for u,v,w in edges:adj[u].append((v,w))
    hits=[0]*n;heap=[(0,0,s)];serial=1;answer=[]
    while heap and len(answer)<k:
        cost,_,u=heappop(heap)
        if hits[u]>=k:continue
        hits[u]+=1
        if u==t:answer.append(cost)
        for v,w in adj[u]:
            if hits[v]<k:heappush(heap,(cost+w,serial,v));serial+=1
    return answer

def path_check(edges,path,s,t):
    u=s;value=0
    for e in path:
        a,v,w=edges[e];check(a==u,'path arc continuity');u=v;value+=w
    check(u==t,'path endpoint');return value

def safe(value):
    if value==inf:return 'INF'
    if isinstance(value,dict):return {str(k):safe(v) for k,v in value.items()}
    if isinstance(value,(tuple,list)):return [safe(v) for v in value]
    return value

def main():
    counts={'Lawler_families':0,'directed_graphs':0,'source_target_pairs':0,'walks_checked':0,'undirected_graphs':0,'spanning_trees_checked':0}
    rng=random.Random(41009)
    # Named interface boundaries, beyond exhaustive small-graph coverage below.
    def forbidden_oracle(prefix):raise AssertionError('k=0 must not query oracle')
    check(lawler(0,forbidden_oracle,0)[0]==[],'Lawler zero request')
    check(lawler(0,lambda prefix:(7,()),2)[0]==[(7,())],'Lawler zero-dimensional solution')
    check(lawler(0,lambda prefix:None,2)[0]==[],'Lawler empty family')
    zero_loop=[(0,1,1),(1,1,0),(1,2,1)]
    zr,zm=eppstein_index(3,zero_loop,0,2,40)
    check(len(zr)==40 and all(x[0]==2 for x in zr),'infinite equal-cost layer finite k')
    zp=[decode_walk(zr,j,zm['tree'],zero_loop,0,2) for j in range(40)]
    check(zp==[(0,)+(1,)*j+(2,) for j in range(40)],'zero-loop identity and explicit length')
    terminal_loop=[(0,0,0)]
    er,em=eppstein_index(1,terminal_loop,0,0,20)
    check([decode_walk(er,j,em['tree'],terminal_loop,0,0) for j in range(20)]==[(0,)*j for j in range(20)],'empty walk and continuations after target')
    check(yen(1,terminal_loop,0,0,20)[0]==[(0,(),(0,))],'simple source equals target')
    check(len(yen(2,[(0,1,0),(0,1,0)],0,1,10)[0])==2,'parallel arcs are distinct paths')
    check(yen(2,[],0,1,5)[0]==[] and eppstein_index(2,[],0,1,5)[0]==[],'unreachable target')
    check(yen(2,[],0,1,0)[0]==[] and eppstein_index(2,[],0,1,0)[0]==[],'zero requested paths')
    check(list(reverse_search_trees(0,[]))==[] and list(reverse_search_trees(1,[]))==[()],'empty versus singleton graph')
    boundary_checks=11
    for bits in range(8):
        universe=list(product((0,1),repeat=bits))
        for _ in range(60):
            feasible={x:rng.randrange(-8,15) for x in universe if rng.random()<.6}
            def oracle(prefix):
                return min(((c,x) for x,c in feasible.items() if x[:len(prefix)]==prefix),default=None)
            out,stats=lawler(bits,oracle,len(feasible)+1)
            check(len(out)==len(feasible) and len({x for c,x in out})==len(out),'Lawler exhaustive completeness')
            check([c for c,x in out]==sorted(feasible.values()),'Lawler costs')
            check(stats['oracle_calls']<=1+len(out)*bits,'Lawler calls');counts['Lawler_families']+=1
    arcs=[(u,v) for u in range(3) for v in range(3) if u!=v]
    graphs=[]
    for options in product([None,0,1],repeat=6):graphs.append((3,[(u,v,w) for (u,v),w in zip(arcs,options) if w is not None]))
    for _ in range(100):
        n=rng.randrange(1,7);edges=[(rng.randrange(n),rng.randrange(n),rng.randrange(4)) for _ in range(rng.randrange(0,2*n+1))];graphs.append((n,edges))
    for n,edges in graphs:
        for s in range(n):
            for t in range(n):
                simple=brute_simple(n,edges,s,t);actual,stats=yen(n,edges,s,t,len(simple)+1)
                check([c for c,p,v in actual]==[c for c,p in simple],'Yen complete costs');check({p for c,p,v in actual}=={p for c,p in simple},'Yen complete distinct paths')
                check(all(len(v)==len(set(v)) for c,p,v in actual),'Yen simple witnesses')
                check(stats['spur_calls']<=1+len(actual)*max(0,n-1),'Yen oracle bound')
                records,meta=eppstein_index(n,edges,s,t,15);oracle=brute_walk_costs(n,edges,s,t,15)
                check([row[0] for row in records]==oracle,('Eppstein cost multiset',n,edges,s,t,records,oracle))
                seen=set()
                for j,row in enumerate(records):
                    path=decode_walk(records,j,meta['tree'],edges,s,t);check(path_check(edges,path,s,t)==row[0],'sidetrack decoding cost');check(path not in seen,'walk duplicate');seen.add(path);counts['walks_checked']+=1
                counts['source_target_pairs']+=1
        counts['directed_graphs']+=1
    for n in range(1,6):
        edges_all=list(combinations(range(n),2))
        for mask in range(1<<len(edges_all)):
            edges=[e for j,e in enumerate(edges_all) if mask>>j&1];expected=[]
            for ids in combinations(range(len(edges)),n-1):
                if root_spanning_tree(n,[edges[i] for i in ids]) is not None:expected.append(ids)
            stats={};actual=list(reverse_search_trees(n,edges,stats))
            check(set(actual)==set(expected) and len(actual)==len(set(actual)),'reverse exhaustive trees')
            check(stats.get('neighbor_slots',0)==len(actual)*len(edges)*(n-1),'each neighbor slot exactly once')
            check(stats.get('backtracks',0)==max(0,len(actual)-1),'one backtrack per nonroot')
            counts['undirected_graphs']+=1;counts['spanning_trees_checked']+=len(actual)
    weights=[1,3,2,4]
    def choose_two(prefix):
        candidates=[x for x in product((0,1),repeat=4) if sum(x)==2 and x[:len(prefix)]==prefix]
        return min(((sum(a*b for a,b in zip(x,weights)),x) for x in candidates),default=None)
    ranked,partition=lawler(4,choose_two,7)
    graph=[(0,1,1),(1,3,2),(0,2,2),(2,3,2),(1,2,1),(2,1,1),(0,3,6),(3,1,2)]
    simple,ystats=yen(4,graph,0,3,20);records,meta=eppstein_index(4,graph,0,3,12)
    walks=[{'cost':row[0],'edges':decode_walk(records,j,meta['tree'],graph,0,3),'implicit':row} for j,row in enumerate(records)]
    undirected=[(0,1),(1,2),(2,3),(3,0),(0,2)];trees=list(reverse_search_trees(4,undirected))
    result={'status':'PASS','counts':counts,'named_boundary_checks':boundary_checks,'examples':{'Lawler':{'weights':weights,'ranked':ranked,'stats':partition},'paths':{'edges':graph,'simple':simple,'Yen_stats':ystats,'walks':walks,'tree':meta['tree'],'distances_to_target':meta['distances'],'deltas':meta['delta']},'reverse_search':{'edges':undirected,'trees':trees}}}
    Path(__file__).with_name('algorithms-kbest-enumeration-results.json').write_text(json.dumps(safe(result),ensure_ascii=False,indent=2)+'\n');print(json.dumps(counts))

if __name__=='__main__':main()
