#!/usr/bin/env python3
"""Exact nonnegative-integer shortest paths; checks remain active under -O."""
from collections import deque
from heapq import heappush, heappop
from itertools import product
from math import inf
from pathlib import Path
import json,random


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


def validate_graph(adj, source, maximum=None):
    n=len(adj)
    if not 0<=source<n: raise ValueError('source out of range')
    for row in adj:
        for v,w in row:
            if type(v) is not int or not 0<=v<n: raise ValueError('bad endpoint')
            if type(w) is not int or w<0 or (maximum is not None and w>maximum):
                raise ValueError('edge weight outside contract')


def zero_one_bfs(adj, source):
    validate_graph(adj,source,1)
    n=len(adj);dist=[inf]*n;parent=[None]*n;dist[source]=0
    queue=deque([(0,source)]);trace=[];scans=0
    while queue:
        check(all(queue[i][0]<=queue[i+1][0] for i in range(len(queue)-1)),'deque snapshot order')
        check(queue[-1][0]<=queue[0][0]+1,'two levels')
        value,u=queue.popleft()
        if value!=dist[u]:
            trace.append(('stale',u,value));continue
        trace.append(('settle',u,value))
        for v,w in adj[u]:
            scans+=1;candidate=value+w
            if candidate<dist[v]:
                dist[v]=candidate;parent[v]=(u,w)
                (queue.appendleft if w==0 else queue.append)((candidate,v))
    check(len(trace)<=2*n,'at most two records per vertex')
    return dist,parent,{'trace':trace,'edge_scans':scans}


def dial(adj,source,maximum):
    if type(maximum) is not int or maximum<0:raise ValueError('bad C')
    validate_graph(adj,source,maximum)
    n=len(adj);dist=[inf]*n;parent=[None]*n;dist[source]=0
    buckets=[deque() for _ in range(maximum+1)]
    buckets[0].append((0,source));pending=1;current=0;scans=0;trace=[];skipped=0
    while pending:
        while not buckets[current%(maximum+1)]:
            current+=1;skipped+=1
        value,u=buckets[current%(maximum+1)].popleft();pending-=1
        check(value==current,('circular bucket alias',value,current))
        if value!=dist[u]:trace.append(('stale',u,value));continue
        trace.append(('settle',u,value))
        for v,w in adj[u]:
            scans+=1;candidate=value+w
            if candidate<dist[v]:
                dist[v]=candidate;parent[v]=(u,w)
                buckets[candidate%(maximum+1)].append((candidate,v));pending+=1
    check(current<=n*maximum,'Dial cursor bound')
    return dist,parent,{'trace':trace,'edge_scans':scans,'cursor':current,'empty_advances':skipped}


class RadixHeap:
    def __init__(self,upper):
        if type(upper) is not int or upper<0:raise ValueError('bad key upper bound')
        self.upper=upper;self.width=upper.bit_length();self.last=0;self.size=0
        self.buckets=[[] for _ in range(self.width+1)]
        self.moves=0;self.scanned_buckets=0;self.insertions=0;self.events=[]
    def push(self,key,item):
        if type(key) is not int or not self.last<=key<=self.upper:raise ValueError('monotone key contract')
        index=(key^self.last).bit_length()
        self.buckets[index].append((key,item));self.size+=1;self.insertions+=1
    def pop(self):
        if not self.size:raise IndexError('empty radix heap')
        if not self.buckets[0]:
            index=1
            while not self.buckets[index]:index+=1
            self.scanned_buckets+=index
            old_last=self.last;self.last=min(key for key,_ in self.buckets[index])
            items=self.buckets[index];self.buckets[index]=[]
            distribution=[]
            for key,item in items:
                new_index=(key^self.last).bit_length();check(new_index<index,'strict bucket descent')
                self.buckets[new_index].append((key,item));self.moves+=1
                distribution.append((key,new_index))
            self.events.append((old_last,self.last,index,distribution))
            for j,row in enumerate(self.buckets):
                check(all((key^self.last).bit_length()==j for key,_ in row),'all bucket invariants')
        self.size-=1
        return self.buckets[0].pop()


def radix_dijkstra(adj,source):
    validate_graph(adj,source)
    n=len(adj);maximum=max((w for row in adj for _,w in row),default=0)
    queue=RadixHeap(n*maximum);dist=[inf]*n;parent=[None]*n;dist[source]=0
    queue.push(0,source);trace=[];scans=0
    while queue.size:
        value,u=queue.pop()
        if value!=dist[u]:trace.append(('stale',u,value));continue
        trace.append(('settle',u,value))
        for v,w in adj[u]:
            scans+=1;candidate=value+w
            if candidate<dist[v]:
                dist[v]=candidate;parent[v]=(u,w);queue.push(candidate,v)
    check(queue.moves<=queue.insertions*queue.width,'radix move budget')
    return dist,parent,{'trace':trace,'edge_scans':scans,'moves':queue.moves,'insertions':queue.insertions,'width':queue.width,'events':queue.events}


def reconstruct(parent,source,target):
    if target!=source and parent[target] is None:return []
    path=[target]
    while path[-1]!=source:
        path.append(parent[path[-1]][0])
        check(len(path)<=len(parent),'parent cycle')
    return path[::-1]


def bidirectional_dijkstra(adj,source,target):
    validate_graph(adj,source)
    n=len(adj)
    if not 0<=target<n:raise ValueError('target out of range')
    if source==target:return 0,[source],{'trace':[],'stop':(0,0,0)}
    reverse=[[] for _ in adj]
    for u,row in enumerate(adj):
        for v,w in row:reverse[v].append((u,w))
    graphs=[adj,reverse];dist=[[inf]*n for _ in range(2)];parents=[[None]*n for _ in range(2)]
    settled=[[False]*n for _ in range(2)]
    dist[0][source]=0;dist[1][target]=0;queues=[[(0,source)],[(0,target)]]
    mu=inf;bridge=None;direction=0;trace=[]
    def clean(which):
        while queues[which] and (queues[which][0][0]!=dist[which][queues[which][0][1]] or settled[which][queues[which][0][1]]):
            heappop(queues[which])
        return queues[which][0][0] if queues[which] else inf
    while True:
        alpha,beta=clean(0),clean(1)
        if alpha==inf or beta==inf or alpha+beta>=mu:break
        value,u=heappop(queues[direction]);settled[direction][u]=True
        other=1-direction
        if settled[other][u] and dist[0][u]+dist[1][u]<mu:
            mu=dist[0][u]+dist[1][u];bridge=(u,u,None)
        for v,w in graphs[direction][u]:
            if settled[other][v]:
                candidate=value+w+dist[other][v]
                if candidate<mu:
                    mu=candidate;bridge=(u,v,w) if direction==0 else (v,u,w)
            candidate=value+w
            if candidate<dist[direction][v]:
                dist[direction][v]=candidate;parents[direction][v]=(u,w)
                heappush(queues[direction],(candidate,v))
        trace.append({'direction':direction,'vertex':u,'settled_distance':value,'mu':mu})
        direction=other
    if bridge is None:return inf,[],{'trace':trace,'stop':(alpha,beta,mu)}
    a,b,w=bridge;path=reconstruct(parents[0],source,a)
    if w is not None:path.append(b)
    while path[-1]!=target:
        path.append(parents[1][path[-1]][0])
        check(len(path)<=2*n,'reverse parent cycle')
    # A zero-cost repeated cycle may occur when two optimal halves meet again; erase it.
    clean_path=[];position={}
    for vertex in path:
        if vertex in position:
            cut=position[vertex]
            while len(clean_path)>cut+1:
                del position[clean_path.pop()]
        else:position[vertex]=len(clean_path);clean_path.append(vertex)
    return mu,clean_path,{'trace':trace,'stop':(alpha,beta,mu),'bridge':bridge}


def floyd_warshall(adj):
    n=len(adj);distance=[[inf]*n for _ in adj]
    for u in range(n):
        distance[u][u]=0
        for v,w in adj[u]:distance[u][v]=min(distance[u][v],w)
    for k in range(n):
        for i in range(n):
            for j in range(n):distance[i][j]=min(distance[i][j],distance[i][k]+distance[k][j])
    return distance


def path_cost(adj,path):
    if not path:return inf
    return sum(min(w for v,w in adj[a] if v==b) for a,b in zip(path,path[1:]))


def verify_graph(adj):
    n=len(adj);maximum=max((w for row in adj for _,w in row),default=0);oracle=floyd_warshall(adj)
    for source in range(n):
        implementations=[dial(adj,source,maximum),radix_dijkstra(adj,source)]
        if maximum<=1:implementations.append(zero_one_bfs(adj,source))
        for dist,parents,stats in implementations:
            check(dist==oracle[source],('distances',adj,source,dist,oracle[source]))
            check(stats['edge_scans']<=sum(map(len,adj)),'each edge scanned once')
            settled=[x[1] for x in stats['trace'] if x[0]=='settle']
            check(len(settled)==len(set(settled)),'each vertex settled once')
            for target in range(n):
                path=reconstruct(parents,source,target)
                check(path_cost(adj,path)==dist[target],('parent path',adj,path,dist[target]))
        for target in range(n):
            value,path,stats=bidirectional_dijkstra(adj,source,target)
            check(value==oracle[source][target] and path_cost(adj,path)==value,('bidirectional',adj,source,target,value,path,oracle[source][target],stats))
    return oracle


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


def main():
    graphs=0;arcs=[(u,v) for u in range(3) for v in range(3) if u!=v]
    for weights in product([None,0,1,3],repeat=len(arcs)):
        adj=[[] for _ in range(3)]
        for (u,v),w in zip(arcs,weights):
            if w is not None:adj[u].append((v,w))
        verify_graph(adj);graphs+=1
    rng=random.Random(314159)
    for _ in range(150):
        n=rng.randrange(1,15);adj=[[] for _ in range(n)]
        for _ in range(rng.randrange(4*n+1)):
            adj[rng.randrange(n)].append((rng.randrange(n),rng.randrange(17)))
        verify_graph(adj);graphs+=1
    heap_operations=0
    for _ in range(100):
        heap=RadixHeap(65535);items={};serial=0
        for _ in range(300):
            if not items or rng.random()<0.64:
                key=rng.randrange(heap.last,65536);heap.push(key,serial);items[serial]=key;serial+=1
            else:
                key,item=heap.pop();check(key==min(items.values()) and items[item]==key,'heap multiset oracle');del items[item]
            heap_operations+=1
        while items:
            key,item=heap.pop();check(key==min(items.values()) and items[item]==key,'heap final oracle');del items[item];heap_operations+=1
        check(heap.moves<=heap.insertions*heap.width,'heap movement bound')
    g01=[[(1,1),(2,0)],[(3,0)],[(1,0),(3,1)],[(4,1)],[],[]]
    integer=[[(1,4),(2,1)],[(3,0),(4,6)],[(1,1),(3,5)],[(4,2)],[],[]]
    meet=[[ (1,4),(2,1)],[(3,4)],[(3,6)],[]]
    examples=[]
    for adj in [g01,integer,meet,[[(0,0)]],[[],[]]]:
        oracle=verify_graph(adj);C=max((w for row in adj for _,w in row),default=0)
        row={'adjacency':adj,'oracle':oracle,'dial':dial(adj,0,C),'radix':radix_dijkstra(adj,0),'bidirectional':bidirectional_dijkstra(adj,0,4 if len(adj)==6 else len(adj)-1)}
        if C<=1:row['zero_one']=zero_one_bfs(adj,0)
        examples.append(row)
    for fn in [lambda:zero_one_bfs([[(0,2)]],0),lambda:dial([[(0,1)]],0,0),lambda:radix_dijkstra([[(0,-1)]],0)]:
        try:fn()
        except ValueError:pass
        else:raise AssertionError('bad edge accepted')
    heap=RadixHeap(10);heap.push(4,'x');heap.pop()
    try:heap.push(3,'bad')
    except ValueError:pass
    else:raise AssertionError('nonmonotone key accepted')
    result={'status':'PASS','counts':{'graphs':graphs,'three_vertex_exhaustive_graphs':4096,'random_multigraphs':150,'heap_operations':heap_operations},'examples':examples}
    Path(__file__).with_name('algorithms-integer-paths-results.json').write_text(json.dumps(safe_json(result),ensure_ascii=False,indent=2)+'\n')
    print(json.dumps(result['counts']))

if __name__=='__main__':main()
