#!/usr/bin/env python3
"""Exact teaching kernels; Python standard library only. No third-party solver.

Arc records are (unique integer ID, tail, head, capacity, cost).
A supplied flow and demand vector are checked before optimization.
Tracing and audit scans are optional and excluded from core operation bounds.
"""
from fractions import Fraction as Q
from collections import deque
from itertools import product
from pathlib import Path
import json
import random


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


def exact(x):
    if isinstance(x, bool) or not isinstance(x, (int, Q)):
        raise ValueError('only exact integers/Fractions are accepted')
    return Q(x)


class Network:
    def __init__(self, n, arcs, flow=None):
        if type(n) is not int or n < 0:
            raise ValueError('nonnegative integer vertex count required')
        self.n = n
        self.ids, self.u, self.v, self.cap, self.cost = [], [], [], [], []
        seen = set()
        for ident, u, v, cap, cost in arcs:
            if type(ident) is not int or ident in seen:
                raise ValueError('distinct integer arc IDs required')
            if type(u) is not int or type(v) is not int or not (0 <= u < n and 0 <= v < n):
                raise ValueError('vertex outside range')
            seen.add(ident)
            cap, cost = exact(cap), exact(cost)
            if cap < 0:
                raise ValueError('negative capacity')
            self.ids.append(ident); self.u.append(u); self.v.append(v)
            self.cap.append(cap); self.cost.append(cost)
        self.m = len(self.ids)
        self.f = [Q(0)] * self.m if flow is None else list(map(exact, flow))
        if len(self.f) != self.m or any(not 0 <= x <= c for x, c in zip(self.f, self.cap)):
            raise ValueError('invalid initial flow')
        self.adj = [[] for _ in range(n)]
        for i in range(self.m):
            self.adj[self.u[i]].append(2*i)
            self.adj[self.v[i]].append(2*i+1)

    def tail(self, a):
        return self.v[a//2] if a & 1 else self.u[a//2]

    def head(self, a):
        return self.u[a//2] if a & 1 else self.v[a//2]

    def residual(self, a):
        i = a//2
        return self.f[i] if a & 1 else self.cap[i] - self.f[i]

    def price_cost(self, a, p):
        c = -self.cost[a//2] if a & 1 else self.cost[a//2]
        return c + p[self.tail(a)] - p[self.head(a)]

    def push(self, a, amount, excess=None):
        check(0 < amount <= self.residual(a), 'invalid residual push')
        self.f[a//2] += -amount if a & 1 else amount
        if excess is not None:
            excess[self.tail(a)] -= amount
            excess[self.head(a)] += amount

    def edges(self):
        return [(a, self.tail(a), self.head(a), -self.cost[a//2] if a & 1 else self.cost[a//2])
                for a in range(2*self.m) if self.residual(a) > 0]

    def balance(self):
        b = [Q(0)] * self.n
        for u, v, f in zip(self.u, self.v, self.f):
            b[u] -= f; b[v] += f
        return b

    def total_cost(self):
        return sum((f*c for f, c in zip(self.f, self.cost)), Q(0))

    def named(self, a):
        return [self.ids[a//2], '-' if a & 1 else '+']


def initialized(n, arcs, flow, demand):
    g = Network(n, arcs, flow)
    b = g.balance() if demand is None else list(map(exact, demand))
    if len(b) != n or g.balance() != b:
        raise ValueError('supplied flow does not meet demand')
    return g, b


def shortest_potential(n, edges):
    """Zero arcs from an implicit supersource; throws on a negative cycle."""
    d = [Q(0)]*n
    if n == 0: return d
    for _ in range(n):
        changed = False
        for _, u, v, c in edges:
            if d[v] > d[u]+c:
                d[v] = d[u]+c; changed = True
        if not changed:
            return d
    raise ValueError('negative cycle: no feasible potential')


def directed_cycle(n, edges):
    """Iterative DFS, returning ordered arc keys (self-loops included)."""
    adj = [[] for _ in range(n)]
    for key, u, v, _ in edges:
        adj[u].append((key, v))
    color, parent = [0]*n, [None]*n
    for root in range(n):
        if color[root]:
            continue
        color[root] = 1
        stack = [(root, 0)]
        while stack:
            u, j = stack[-1]
            if j == len(adj[u]):
                color[u] = 2; stack.pop(); continue
            key, v = adj[u][j]
            stack[-1] = (u, j+1)
            if color[v] == 0:
                parent[v] = (u, key); color[v] = 1; stack.append((v, 0))
            elif color[v] == 1:
                path, x = [], u
                while x != v:
                    x, arc = parent[x]; path.append(arc)
                path.reverse(); path.append(key)
                return path
    return None


def minimum_mean_cycle(n, edges, keep_table=False):
    """All-start exact-length Karp DP, followed by shifted tight-edge DFS."""
    edges = [(key, u, v, exact(c)) for key, u, v, c in edges]
    if type(n) is not int or n < 0 or any(not 0 <= u < n or not 0 <= v < n for _,u,v,_ in edges):
        raise ValueError('invalid graph')
    if len({e[0] for e in edges}) != len(edges):
        raise ValueError('duplicate edge key')
    # O(n+m) vertex compaction, preserving original vertex and arc order.
    used = [False]*n
    for _, u, v, _ in edges:
        used[u] = used[v] = True
    vertices = [v for v in range(n) if used[v]]
    rank = [-1]*n
    for i, v in enumerate(vertices): rank[v] = i
    es = [(key, rank[u], rank[v], c) for key,u,v,c in edges]
    N = len(vertices)
    dp = [[Q(0)]*N]
    for _ in range(N):
        row = [None]*N
        for _, u, v, c in es:
            x = dp[-1][u]
            if x is not None and (row[v] is None or x+c < row[v]):
                row[v] = x+c
        dp.append(row)
    best = None
    for v in range(N):
        if dp[N][v] is not None:
            value = max((dp[N][v]-dp[k][v]) / (N-k)
                        for k in range(N) if dp[k][v] is not None)
            if best is None or value < best:
                best = value
    if best is None:
        potential = shortest_potential(n, edges)
        return {'mean': None, 'cycle': None, 'potential': potential,
                'vertices': vertices, 'table': dp if keep_table else None}
    shifted = [(key, u, v, c-best) for key,u,v,c in es]
    p = shortest_potential(N, shifted)
    tight = [e for e in shifted if e[3]+p[e[1]]-p[e[2]] == 0]
    cycle = directed_cycle(N, tight)
    check(cycle is not None, 'minimum-mean tight graph must contain a cycle')
    potential = [Q(0)]*n
    for i, v in enumerate(vertices): potential[v] = p[i]
    return {'mean': best, 'cycle': cycle, 'potential': potential,
            'vertices': vertices, 'table': dp if keep_table else None}


def verify_mean(n, edges, answer):
    index = {e[0]: e for e in edges}
    if answer['mean'] is None:
        check(directed_cycle(n, edges) is None, 'false acyclic result')
        return
    keys, mu, p = answer['cycle'], answer['mean'], answer['potential']
    check(keys and len(set(keys)) == len(keys), 'empty or repeated cycle arc')
    es = [index[k] for k in keys]
    check(len({e[1] for e in es}) == len(es), 'cycle not simple')
    check(all(es[i][2] == es[(i+1)%len(es)][1] for i in range(len(es))), 'broken cycle')
    check(sum(e[3] for e in es) == mu*len(es), 'cycle mean mismatch')
    check(all(c+p[u]-p[v] >= mu for _,u,v,c in edges), 'mean lower certificate')


def verify_optimum(n, arcs, flow, demand, potential):
    g, _ = initialized(n, arcs, flow, demand)
    check(len(potential) == n, 'potential dimension')
    check(all(g.price_cost(a, potential) >= 0 for a in range(2*g.m)
              if g.residual(a) > 0), 'negative reduced residual cost')
    return g.total_cost()


def mean_cycle_canceling(n, arcs, flow=None, demand=None, trace=False):
    g, b = initialized(n, arcs, flow, demand)
    log, iterations = [], 0
    while True:
        result = minimum_mean_cycle(n, g.edges())
        mu = result['mean']
        if mu is None or mu >= 0:
            p = result['potential']
            break
        cycle = result['cycle']
        delta = min(g.residual(a) for a in cycle)
        before = g.total_cost() if trace else None
        for a in cycle: g.push(a, delta)
        iterations += 1
        if trace:
            log.append({'mean': mu, 'cycle': [g.named(a) for a in cycle],
                        'amount': delta, 'cost_before': before, 'cost_after': g.total_cost(),
                        'flow': g.f.copy(), 'potential_before': result['potential']})
    verify_optimum(n, arcs, g.f, b, p)
    return {'flow': g.f, 'potential': p, 'cost': g.total_cost(),
            'iterations': iterations, 'trace': log}


def repair_path(g, p, excess, source):
    """Lazy binary-heap Dijkstra, stop at first settled deficit."""
    from heapq import heappush, heappop
    n = g.n
    d, pred, settled = [None]*n, [None]*n, [False]*n
    d[source] = Q(0)
    queue = [(Q(0), source)]
    target = None
    scans = 0
    while queue:
        value, u = heappop(queue)
        if settled[u] or d[u] != value: continue
        settled[u] = True
        if excess[u] < 0:
            target = u; break
        for a in g.adj[u]:
            scans += 1
            if g.residual(a) <= 0: continue
            v, c = g.head(a), g.price_cost(a, p)
            check(c >= 0, 'Dijkstra needs globally nonnegative reduced costs')
            nv = value+c
            if d[v] is None or nv < d[v]:
                d[v], pred[v] = nv, a; heappush(queue, (nv, v))
    check(target is not None, 'known feasible circulation implies reachable deficit')
    D = d[target]
    for v in range(n): p[v] += D if d[v] is None else min(d[v], D)
    path, v = [], target
    while v != source:
        a = pred[v]; check(a is not None, 'missing predecessor')
        path.append(a); v = g.tail(a)
    path.reverse()
    delta = min(excess[source], -excess[target], *(g.residual(a) for a in path))
    for a in path: g.push(a, delta, excess)
    return path, delta, D, scans


def capacity_scaling_circulation(n, arcs, trace=False):
    g = Network(n, arcs)
    if any(c.denominator != 1 for c in g.cap):
        raise ValueError('capacity-bit scaling requires integer capacities')
    original = [int(c) for c in g.cap]
    bits = max(original, default=0).bit_length()
    g.cap = [Q(0)]*g.m
    p, log = [Q(0)]*n, []
    augmentations = 0
    for bit in range(bits-1, -1, -1):
        g.f = [2*f for f in g.f]
        g.cap = [Q(2*c+((u >> bit)&1)) for c,u in zip(g.cap, original)]
        excess, saturated = [Q(0)]*n, []
        for a in range(2*g.m):
            if g.residual(a) > 0 and g.price_cost(a,p) < 0:
                delta = g.residual(a)
                check(delta == 1, 'only a new unit can violate old dual feasibility')
                g.push(a, delta, excess)
                if trace: saturated.append(g.named(a))
        initial_excess = sum(max(x,0) for x in excess)
        repairs = []
        # Each repair reduces integral total positive excess by >= 1.
        active = deque(v for v in range(n) if excess[v] > 0)
        while active:
            s = active[0]
            if excess[s] <= 0: active.popleft(); continue
            path, delta, distance, scans = repair_path(g,p,excess,s)
            augmentations += 1
            if trace:
                repairs.append({'path':[g.named(a) for a in path], 'amount':delta,
                                'reduced_distance':distance, 'excess':excess.copy(),
                                'potential':p.copy(), 'arc_scans':scans})
        check(all(x == 0 for x in excess), 'unfinished balance repair')
        if trace:
            log.append({'bit':bit, 'capacity':g.cap.copy(), 'negative_units':saturated,
                        'initial_total_excess':initial_excess, 'repairs':repairs,
                        'flow':g.f.copy(), 'potential':p.copy(), 'cost':g.total_cost()})
    verify_optimum(n, arcs, g.f, [Q(0)]*n, p)
    return {'flow':g.f, 'potential':p, 'cost':g.total_cost(),
            'bits':bits, 'augmentations':augmentations, 'trace':log}


def capacity_scaling(n, arcs, flow=None, demand=None, trace=False):
    base, b = initialized(n, arcs, flow, demand)
    if any(x.denominator != 1 for x in base.cap+base.f):
        raise ValueError('capacity scaling requires integer capacities and initial flow')
    residual = [(a,u,v,base.residual(a),c) for a,u,v,c in base.edges()]
    result = capacity_scaling_circulation(n, residual, trace)
    final = base.f.copy()
    for (a,_,_,_,_), amount in zip(residual,result['flow']):
        final[a//2] += -amount if a & 1 else amount
    p = result['potential']
    # Every final original residual direction occurs either as an auxiliary
    # forward record or as the reverse of its opposite auxiliary record.
    verify_optimum(n,arcs,final,b,p)
    return {**result, 'auxiliary_flow':result['flow'], 'auxiliary_arcs':residual,
            'flow':final, 'potential':p, 'cost':sum(f*c for f,c in zip(final,base.cost))}


def cost_scaling(n, arcs, flow=None, demand=None, trace=False, audit=False):
    g, b = initialized(n, arcs, flow, demand)
    if any(c.denominator != 1 for c in g.cost):
        raise ValueError('epsilon scaling requires integer costs')
    C = max((abs(int(c)) for c in g.cost), default=0)
    epsilon = Q(1 << (C-1).bit_length()) if C else Q(0)
    p, logs = [Q(0)]*n, []
    totals = {'phases':0, 'pushes':0, 'relabels':0, 'arc_tests':0}
    while n*epsilon >= 1:
        epsilon /= 2
        excess, events = [Q(0)]*n, []
        phase = {'epsilon':epsilon,'pushes':0,'relabels':0,'arc_tests':0}
        start_p = p.copy() if audit else None
        relabel_count = [0]*n if audit else None
        saturated = []
        for a in range(2*g.m):
            if g.residual(a) > 0 and g.price_cost(a,p) < 0:
                amount = g.residual(a); g.push(a,amount,excess)
                if trace: saturated.append({'arc':g.named(a),'amount':amount})
        initial_excess = excess.copy() if trace else None
        current, queued = [0]*n, [False]*n
        queue = deque()
        for v in range(n):
            if excess[v] > 0: queue.append(v); queued[v] = True
        while queue:
            u = queue.popleft(); queued[u] = False
            while excess[u] > 0:
                if current[u] == len(g.adj[u]):
                    candidates = [p[g.head(a)] - (-g.cost[a//2] if a&1 else g.cost[a//2]) - epsilon
                                  for a in g.adj[u] if g.head(a) != u and g.residual(a) > 0]
                    check(candidates, 'active vertex has a residual path to a deficit')
                    old = p[u]; p[u] = max(candidates); current[u] = 0
                    check(p[u] <= old-epsilon, 'price must decrease by at least epsilon')
                    phase['relabels'] += 1
                    if trace: events.append({'relabel':u,'old':old,'new':p[u]})
                    if audit:
                        relabel_count[u] += 1
                        check(start_p[u]-p[u] <= 3*(n-1)*epsilon, 'phase price-drop bound')
                        check(relabel_count[u] <= 3*(n-1), 'phase relabel bound')
                else:
                    a = g.adj[u][current[u]]; phase['arc_tests'] += 1
                    if g.head(a) == u or g.residual(a) <= 0 or g.price_cost(a,p) >= 0:
                        current[u] += 1; continue
                    v = g.head(a)
                    amount = min(excess[u],g.residual(a)); old_excess = excess[v]
                    g.push(a,amount,excess); phase['pushes'] += 1
                    if excess[v] > 0 and old_excess <= 0 and not queued[v] and v != u:
                        queue.append(v); queued[v] = True
                    if trace: events.append({'push':g.named(a),'amount':amount,'excess':excess.copy()})
                if audit:
                    check(all(g.price_cost(a,p) >= -epsilon for a in range(2*g.m)
                              if g.residual(a) > 0), 'epsilon invariant')
                    admissible = [(a,g.tail(a),g.head(a),0) for a in range(2*g.m)
                                  if g.residual(a)>0 and g.price_cost(a,p)<0]
                    check(directed_cycle(n,admissible) is None, 'admissible graph must be acyclic')
        check(g.balance() == b, 'phase output must be feasible')
        totals['phases'] += 1
        for key in ('pushes','relabels','arc_tests'): totals[key] += phase[key]
        if trace:
            logs.append({**phase,'saturated':saturated,'initial_excess':initial_excess,
                         'events':events,'flow':g.f.copy(),'potential':p.copy(),'cost':g.total_cost()})
    epsilon_p = p.copy()
    p = shortest_potential(n,g.edges())
    verify_optimum(n,arcs,g.f,b,p)
    return {'flow':g.f,'potential':p,'epsilon_potential':epsilon_p,'epsilon':epsilon,
            'cost':g.total_cost(),'stats':totals,'trace':logs}


def brute_flows(n, arcs, demand):
    g = Network(n,arcs)
    check(all(c.denominator==1 for c in g.cap),'integer oracle only')
    best, flows = None, []
    for f in product(*(range(int(c)+1) for c in g.cap)):
        g.f = list(map(Q,f))
        if g.balance() != demand: continue
        value = g.total_cost()
        if best is None or value < best: best,flows = value,[f]
        elif value == best: flows.append(f)
    return best,flows


def brute_cycles(n, edges):
    adj = [[] for _ in range(n)]
    for e in edges: adj[e[1]].append(e)
    best = None
    for start in range(n):
        todo = [(start,{start},Q(0),0)]
        while todo:
            v,seen,cost,length = todo.pop()
            for _,_,w,c in adj[v]:
                if w==start:
                    value=(cost+c)/(length+1)
                    if best is None or value<best: best=value
                elif w>start and w not in seen:
                    todo.append((w,seen|{w},cost+c,length+1))
    return best


def jsonable(x):
    if isinstance(x,Q): return x.numerator if x.denominator==1 else str(x)
    if isinstance(x,dict): return {str(k):jsonable(v) for k,v in x.items()}
    if isinstance(x,(list,tuple)): return [jsonable(v) for v in x]
    return x


def main():
    arcs = [(10,0,2,3,7),(20,0,1,2,1),(21,0,1,2,2),(30,1,2,4,1),
            (40,1,0,1,0),(50,3,4,2,-4),(51,4,3,3,1),(60,4,4,2,-2),(70,2,1,1,3)]
    initial=[3,0,0,0,0,0,0,0,0]
    demand=[Q(-3),Q(0),Q(3),Q(0),Q(0),Q(0)]
    example={'arcs':arcs,'initial':initial,'demand':demand}
    for name,solver in [('mean',mean_cycle_canceling),('capacity',capacity_scaling),('cost',cost_scaling)]:
        kwargs={'trace':True}
        if name=='cost':kwargs['audit']=True
        example[name]=solver(6,arcs,initial,demand,**kwargs)
        check(example[name]['cost']==-3,'main example optimum')
    edges=[(10,0,1,2),(20,1,2,-5),(30,2,0,0),(40,2,3,4),(50,3,3,-Q(1,2))]
    mean_example=minimum_mean_cycle(5,edges,True);verify_mean(5,edges,mean_example)
    rng=random.Random(20261009)
    cases=0;cycle_cases=0
    for n in range(1,5):
        for _ in range(180):
            m=rng.randrange(0,8)
            ar=[(10+7*i,rng.randrange(n),rng.randrange(n),rng.randrange(3),rng.randrange(-5,6)) for i in range(m)]
            f=[rng.randrange(cap+1) for _,_,_,cap,_ in ar]
            b=Network(n,ar,f).balance()
            best,_=brute_flows(n,ar,b)
            for solver in (mean_cycle_canceling,capacity_scaling,cost_scaling):
                got=solver(n,ar,f,b,**({'audit':True} if solver is cost_scaling else {}))
                check(got['cost']==best,'flow enumeration mismatch')
            cases+=1
            es=[(ident,u,v,Q(c)) for ident,u,v,_,c in ar]
            got=minimum_mean_cycle(n,es);verify_mean(n,es,got)
            check(got['mean']==brute_cycles(n,es),'cycle enumeration mismatch');cycle_cases+=1
    rational_cases=0
    for _ in range(120):
        n=4;m=6
        ar=[(i,rng.randrange(n),rng.randrange(n),Q(rng.randrange(4),3),rng.randrange(-6,7)) for i in range(m)]
        f=[Q(rng.randrange(int(3*c)+1),3) for _,_,_,c,_ in ar]
        b=Network(n,ar,f).balance()
        scaled=[(i,u,v,3*c,cost) for i,u,v,c,cost in ar]
        optimum,_=brute_flows(n,scaled,[3*x for x in b])
        for solver in (mean_cycle_canceling,cost_scaling):
            got=solver(n,ar,f,b,**({'audit':True} if solver is cost_scaling else {}))
            check(3*got['cost']==optimum,'rational capacity oracle')
        rational_cases+=1
    for solver in (mean_cycle_canceling,capacity_scaling,cost_scaling):
        check(solver(0,[])['cost']==0,'empty graph')
    epsilon_boundary={'arcs':[(1,0,1,1,-1),(2,1,0,1,0)],'flow':[0,0],
                      'epsilon':Q(1,2),'potential':[0,Q(-1,2)],'negative_cycle_cost':-1}
    g=Network(2,epsilon_boundary['arcs'])
    check(all(g.price_cost(a,epsilon_boundary['potential'])>=-Q(1,2) for a in range(4) if g.residual(a)>0),'boundary epsilon certificate')
    result={'example':example,'mean_example':{'edges':edges,**mean_example},
            'epsilon_boundary':epsilon_boundary,'self_checks':{'integer_flow_instances':cases,
            'directed_cycle_instances':cycle_cases,'rational_capacity_instances':rational_cases,
            'normal_and_optimized_use_explicit_checks':True}}
    path=Path(__file__).with_name('algorithms-flow-scaling-results.json')
    path.write_text(json.dumps(jsonable(result),ensure_ascii=False,indent=2)+'\n')
    print(json.dumps(jsonable({'status':'PASS',**result['self_checks'],
                              'main_cost':example['cost']['cost']}),sort_keys=True))


if __name__=='__main__':main()
