#!/usr/bin/env python3
"""Source flow-to-table bijection plus bounded exact arc-vector diagnostics.

Implements the selected paper reduction, not its FPRAS or approximate sampler.
Negative lower bounds, parallel arcs and loops are explicit. One state is an
integer arc-value vector; path/cycle decompositions are never counted.
"""
import argparse
import json
from collections import Counter
from fractions import Fraction
from pathlib import Path

from contingency_reference import TableSpace, BudgetExceeded, MAX_COUNT_WORK, MAX_STATES

MAX_VERTICES = 16
MAX_ARCS = 16
MAX_INTEGER_BITS = 256
MAX_FLOW_DIAGNOSTICS = 2_000
MAX_ENUM_WORK = 4_000_000


def integer(value):
    if type(value) is not int or value.bit_length() > MAX_INTEGER_BITS:
        raise ValueError('Require signed integers of at most 256 magnitude bits')
    return value


def identifiers(value, maximum, name):
    if not isinstance(value, list) or len(value) > maximum or any(not isinstance(v, str) or not 1 <= len(v) <= 64 for v in value) or len(set(value)) != len(value):
        raise ValueError('Invalid ' + name)
    return value


class FlowModel:
    def __init__(self, record):
        if not isinstance(record, dict) or set(record) != {'vertices', 'arcs', 'balances'}:
            raise ValueError('Require vertices, arcs and balances')
        self.vertices = identifiers(record['vertices'], MAX_VERTICES, 'vertices')[:]
        if not isinstance(record['balances'], dict) or set(record['balances']) != set(self.vertices):
            raise ValueError('Supply one balance for every vertex')
        self.balances = {v: integer(record['balances'][v]) for v in self.vertices}
        if not isinstance(record['arcs'], list) or len(record['arcs']) > MAX_ARCS:
            raise ValueError('At most sixteen explicitly identified arcs')
        self.arcs = []
        for arc in record['arcs']:
            if not isinstance(arc, dict) or set(arc) != {'id', 'tail', 'head', 'lower', 'upper'}:
                raise ValueError('Invalid arc fields')
            identifiers([arc['id']], 1, 'arc ID')
            if not isinstance(arc['tail'], str) or not isinstance(arc['head'], str) or arc['tail'] not in self.balances or arc['head'] not in self.balances:
                raise ValueError('Unknown arc endpoint')
            self.arcs.append(dict(arc, lower=integer(arc['lower']), upper=integer(arc['upper'])))
        identifiers([a['id'] for a in self.arcs], MAX_ARCS, 'arc IDs')
        self.precondition_empty = 'Arc lower exceeds upper' if any(a['lower'] > a['upper'] for a in self.arcs) else (
            'Vertex balances do not sum to zero' if sum(self.balances.values()) else None)
        self.q = len(self.vertices) + len(self.arcs)
        self.positions = {v: i for i, v in enumerate(self.vertices)}
        self.shifted = self.balances.copy()
        for arc in self.arcs:
            self.shifted[arc['tail']] -= arc['lower']
            self.shifted[arc['head']] += arc['lower']
        self.caps = [a['upper'] - a['lower'] for a in self.arcs]
        self.K = None if self.precondition_empty else 1 + 2 * sum(self.caps) + sum(abs(v) for v in self.shifted.values())
        self.rows, self.columns, self.bounds = [], [], []
        if not self.precondition_empty:
            beta = [self.shifted[v] for v in self.vertices] + [0] * len(self.arcs)
            self.rows = [self.K + max(b, 0) for b in beta]
            self.columns = [self.K + max(-b, 0) for b in beta]
            self.bounds = [[self.K if i == j else 0 for j in range(self.q)] for i in range(self.q)]
            for i, arc in enumerate(self.arcs):
                private = len(self.vertices) + i
                self.bounds[self.positions[arc['tail']]][private] = self.caps[i]
                self.bounds[private][self.positions[arc['head']]] = self.caps[i]

    def validate_flow(self, vector):
        if not isinstance(vector, list) or len(vector) != len(self.arcs):
            raise ValueError('One value per arc in declared order')
        balance = {v: 0 for v in self.vertices}
        for arc, value in zip(self.arcs, vector):
            integer(value)
            if not arc['lower'] <= value <= arc['upper']:
                raise ValueError('Arc value violates bounds')
            balance[arc['tail']] += value
            balance[arc['head']] -= value
        if balance != self.balances:
            raise ValueError('Flow balance mismatch')

    def to_table(self, vector):
        self.validate_flow(vector)
        if self.precondition_empty:
            raise ValueError('Empty flow precondition')
        table = [[0] * self.q for _ in range(self.q)]
        for i, (arc, value) in enumerate(zip(self.arcs, vector)):
            private = len(self.vertices) + i
            x = value - arc['lower']
            table[self.positions[arc['tail']]][private] = x
            table[private][self.positions[arc['head']]] = x
        for i in range(self.q):
            table[i][i] = self.rows[i] - sum(table[i])
        self.validate_table(table)
        return table

    def validate_table(self, table):
        if self.precondition_empty or not isinstance(table, list) or len(table) != self.q or any(not isinstance(row, list) or len(row) != self.q for row in table):
            raise ValueError('Invalid reduced table shape/precondition')
        for i, row in enumerate(table):
            for j, value in enumerate(row):
                if type(value) is not int or not 0 <= value <= self.bounds[i][j]:
                    raise ValueError('Reduced table bound violation')
        if [sum(row) for row in table] != self.rows or [sum(table[i][j] for i in range(self.q)) for j in range(self.q)] != self.columns:
            raise ValueError('Reduced table margins differ')

    def from_table(self, table):
        self.validate_table(table)
        vector = []
        for i, arc in enumerate(self.arcs):
            private = len(self.vertices) + i
            first = table[self.positions[arc['tail']]][private]
            second = table[private][self.positions[arc['head']]]
            if first != second:
                raise ValueError('Private-vertex arc copies differ')
            vector.append(arc['lower'] + first)
        self.validate_flow(vector)
        return vector

    def table_record(self):
        return dict(rows=self.rows, columns=self.columns, cell_bounds=self.bounds)


def calculate(record, state_budget=MAX_STATES, count_work_budget=MAX_COUNT_WORK, enum_work_budget=MAX_ENUM_WORK):
    for value, maximum in ((state_budget, MAX_STATES), (count_work_budget, MAX_COUNT_WORK)):
        if type(value) is not int or not 1 <= value <= maximum:
            raise ValueError('Invalid counting budget')
    if type(enum_work_budget) is not int or not 1 <= enum_work_budget <= MAX_ENUM_WORK:
        raise ValueError('Invalid diagnostic work budget')
    model = FlowModel(record)
    result = dict(status='reduction_complete_count_unknown', exact_count=None, feasible=None, flows=None, arc_diagnostics=None,
                  arc_order=[a['id'] for a in model.arcs],
                  reduction=dict(table=model.table_record() if not model.precondition_empty else None,
                                 K=model.K, dimension=model.q, total=sum(model.rows),
                                 index_order=[dict(kind='original_vertex', id=v) for v in model.vertices] + [dict(kind='private_arc_vertex', id=a['id']) for a in model.arcs],
                                 balance_convention='Outflow minus inflow', lower_bound_shifted_balances=model.shifted,
                                 preserves='One table for every original integer arc-value vector, including distinct parallel arcs and loops'),
                  limits=dict(vertices=MAX_VERTICES, arcs=MAX_ARCS, magnitude_bits=MAX_INTEGER_BITS,
                              exact_table_dimension=6, exact_table_total=128, diagnostics=MAX_FLOW_DIAGNOSTICS,
                              state_budget=state_budget, count_work_budget=count_work_budget, enum_work_budget=enum_work_budget),
                  source_connection=dict(family='115', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                                         paper_section='An FPRAS for Cell-Bounded Contingency Tables, bounded integral flows',
                                         source_FPRAS_or_approximate_sampler_implemented=False,
                                         selected_comparator_contains_separate_flow_statement=False),
                  source_proof_verification='not_run', formal_program_verification='not_run',
                  interpretation='Finite reduction/count diagnostics. A state is an arc vector, not a path decomposition. Uniform feasible flows describe a chosen scenario law, not real failure probabilities, route optimality or a multicommodity model. Diagnostic/budget exhaustion remains unknown, never zero.')
    if model.precondition_empty:
        return dict(result, status='empty', exact_count=0, feasible=False, reason=model.precondition_empty, flows=[])
    if model.q == 0:
        return dict(result, status='complete', exact_count=1, feasible=True, flows=[dict(rank=0, vector=[])], arc_diagnostics=[])
    if model.q > 6 or sum(model.rows) > 128:
        return dict(result, reason='Valid reduction exceeds the finite table-reference domain; count and feasibility unknown')
    space = TableSpace(model.table_record(), state_budget, count_work_budget)
    try:
        count = space.count()
    except BudgetExceeded as exc:
        return dict(result, reason=str(exc), count_work=space.work, memo_states=space.states)
    count_work = space.work
    result.update(exact_count=count, feasible=count > 0, count_work=count_work, memo_states=space.states)
    if count == 0:
        return dict(result, status='empty', flows=[])
    if count > MAX_FLOW_DIAGNOSTICS:
        return dict(result, status='count_complete_diagnostics_unknown', reason='Exact flow count exceeds diagnostic enumeration cap')
    flows = []
    space.work, space.work_limit = 0, enum_work_budget
    try:
        for rank in range(count):
            vector = model.from_table(space.unrank(rank))
            space.tick()
            flows.append(dict(rank=rank, vector=vector))
    except BudgetExceeded as exc:
        return dict(result, status='count_complete_diagnostics_unknown', reason=str(exc), diagnostic_work=space.work)
    diagnostics = []
    for i, arc in enumerate(model.arcs):
        values = Counter(f['vector'][i] for f in flows)
        diagnostics.append(dict(arc=arc['id'], minimum=min(values), maximum=max(values),
                                exact_uniform_mean=str(Fraction(sum(v * k for v, k in values.items()), count)),
                                zero_flow_event_probability=str(Fraction(values.get(0, 0), count)),
                                upper_bound_event_probability=str(Fraction(values.get(arc['upper'], 0), count)),
                                histogram=[dict(value=v, count=k, uniform_probability=str(Fraction(k, count))) for v, k in sorted(values.items())]))
    return dict(result, status='complete', flows=flows, arc_diagnostics=diagnostics, diagnostic_work=space.work)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', type=Path)
    args = parser.parse_args()
    body = json.dumps(calculate(json.loads(args.input.read_text())), indent=2) + '\n'
    if args.output:
        with args.output.open('x') as stream:
            stream.write(body)
    else:
        print(body, end='')


if __name__ == '__main__':
    main()
