#!/usr/bin/env python3
"""Exact small mass-action model checks and supplied finite-bound evidence."""
import argparse
import hashlib
import json
import math
import re
from fractions import Fraction as Q
from pathlib import Path

MAX_BYTES = 1048576
MAX_BITS = 8192


def integer(x, a, b):
    if type(x) is not int or not a <= x <= b:
        raise ValueError(f'expected integer in [{a},{b}]')
    return x


def rational(x):
    if type(x) is int:
        q = Q(x)
    elif type(x) is str and len(x) <= 64 and re.fullmatch(r'-?(0|[1-9][0-9]*)(/[1-9][0-9]*)?', x):
        q = Q(x)
    else:
        raise ValueError('use integer or exact integer/fraction string')
    if max(abs(q.numerator).bit_length(), q.denominator.bit_length()) > 32:
        raise ValueError('input rational component exceeds 32 bits')
    return q


def bounded(q):
    if max(abs(q.numerator).bit_length(), q.denominator.bit_length()) > MAX_BITS:
        raise ArithmeticError('intermediate rational bit limit')
    return q


def dot(a, b):
    result = Q(0)
    for x, y in zip(a, b):
        result = bounded(result+bounded(x*y))
    return result


def vector(raw, n, nonnegative=True):
    if type(raw) is not list or len(raw) != n:
        raise ValueError('vector length does not match species')
    result = [rational(x) for x in raw]
    if nonnegative and any(x < 0 for x in result):
        raise ValueError('expected nonnegative vector')
    return result


def parse(data):
    required = {'schema','dynamics','species','complexes','reactions','initial','rate_convention'}
    if type(data) is not dict or not required <= data.keys() or data.keys()-required-{'conservation_vectors','box'}:
        raise ValueError('unexpected or missing model fields')
    if data['schema'] != 'reaction-network-v1' or data['dynamics'] != 'deterministic_mass_action_ode':
        raise ValueError('only the explicit deterministic mass-action ODE schema is supported')
    if data['rate_convention'] not in ('monomial_coefficient','factorial_scaled_ode'):
        raise ValueError('rate convention must be explicit')
    names = data['species']
    if type(names) is not list or not 1 <= len(names) <= 8 or any(type(s) is not str or not re.fullmatch(r'[A-Za-z][A-Za-z0-9_]{0,31}', s) for s in names) or len(set(names)) != len(names):
        raise ValueError('use one to eight distinct simple species identifiers')
    n = len(names); complexes = data['complexes']
    if type(complexes) is not list or not 1 <= len(complexes) <= 64:
        raise ValueError('use one to sixty-four explicit complexes')
    for c in complexes:
        if type(c) is not list or len(c) != n:
            raise ValueError('complex dimension mismatch')
        for x in c: integer(x, 0, 8)
        if sum(c) > 8: raise ValueError('total complex degree exceeds eight')
    if len(set(map(tuple, complexes))) != len(complexes):
        raise ValueError('complex vectors must be unique')
    reactions = data['reactions']
    if type(reactions) is not list or len(reactions) > 128:
        raise ValueError('use at most 128 reactions')
    normalized = {}; removed = []; coefficients = []
    for i, rx in enumerate(reactions):
        if type(rx) is not dict or set(rx) != {'source','target','rate'}:
            raise ValueError('reaction requires exactly source, target, rate')
        source = integer(rx['source'], 0, len(complexes)-1)
        target = integer(rx['target'], 0, len(complexes)-1)
        k = rational(rx['rate'])
        if k < 0: raise ValueError('negative rate is invalid')
        divisor = math.prod(math.factorial(x) for x in complexes[source]) if data['rate_convention'] == 'factorial_scaled_ode' else 1
        coefficient = bounded(k/divisor)
        coefficients.append(dict(index=i, coefficient=str(coefficient), factorial_divisor=divisor))
        if coefficient == 0 or source == target:
            removed.append(dict(index=i, reason='zero_rate' if coefficient == 0 else 'zero_stoichiometric_change'))
        else:
            normalized[source,target] = bounded(normalized.get((source,target), Q(0))+coefficient)
    initial = vector(data['initial'], n)
    supplied = data.get('conservation_vectors', [])
    if type(supplied) is not list or len(supplied) > 8:
        raise ValueError('at most eight conservation vectors')
    conserved = [vector(v, n) for v in supplied]
    box = None
    if 'box' in data:
        if type(data['box']) is not dict or set(data['box']) != {'lower','upper'}:
            raise ValueError('box requires exactly lower and upper')
        lower = vector(data['box']['lower'], n); upper = vector(data['box']['upper'], n)
        if any(l <= 0 or u < l for l, u in zip(lower, upper)):
            raise ValueError('box bounds must be positive and ordered')
        box = (lower, upper)
    return names, complexes, normalized, initial, conserved, box, coefficients, removed


def graph_info(count, edges):
    reach = [[i == j for j in range(count)] for i in range(count)]
    neighbors = [set() for _ in range(count)]
    for i, j in edges:
        reach[i][j] = True; neighbors[i].add(j); neighbors[j].add(i)
    for k in range(count):
        for i in range(count):
            if reach[i][k]:
                for j in range(count): reach[i][j] |= reach[k][j]
    remaining = set(range(count)); scc = []
    while remaining:
        i = min(remaining); component = sorted(j for j in remaining if reach[i][j] and reach[j][i])
        scc.append(component); remaining.difference_update(component)
    remaining = set(range(count)); links = []
    while remaining:
        pending = [min(remaining)]; component = set()
        while pending:
            i = pending.pop()
            if i in component: continue
            component.add(i); pending.extend(neighbors[i]-component)
        links.append(sorted(component)); remaining.difference_update(component)
    failed = [[i,j] for i,j in edges if not reach[j][i]]
    return dict(weakly_reversible=not failed, reversible=all((j,i) in edges for i,j in edges),
                edges_without_return_path=failed, strong_components=scc, linkage_classes=links)


def nullspace(rows, n):
    matrix = [[Q(x) for x in row] for row in rows]
    pivots = []; r = 0
    for column in range(n):
        pivot = next((i for i in range(r, len(matrix)) if matrix[i][column]), None)
        if pivot is None: continue
        matrix[r], matrix[pivot] = matrix[pivot], matrix[r]
        factor = matrix[r][column]
        matrix[r] = [bounded(x/factor) for x in matrix[r]]
        for i in range(len(matrix)):
            if i != r and matrix[i][column]:
                factor = matrix[i][column]
                matrix[i] = [bounded(x-factor*y) for x,y in zip(matrix[i], matrix[r])]
        pivots.append(column); r += 1
    basis = []
    for free in range(n):
        if free in pivots: continue
        v = [Q(0)]*n; v[free] = Q(1)
        for i, p in enumerate(pivots): v[p] = -matrix[i][free]
        basis.append(v)
    return len(pivots), basis


def polynomials(complexes, rates, n):
    result = [{} for _ in range(n)]
    for (source, target), k in rates.items():
        exponent = tuple(complexes[source])
        for i in range(n):
            term = bounded(k*(complexes[target][i]-complexes[source][i]))
            result[i][exponent] = bounded(result[i].get(exponent, Q(0))+term)
    return [{e:c for e,c in p.items() if c} for p in result]


def evaluate(polys, point):
    output = []
    for polynomial in polys:
        total = Q(0)
        for exponent, coefficient in polynomial.items():
            term = coefficient
            for x, power in zip(point, exponent): term = bounded(term*bounded(x**power))
            total = bounded(total+term)
        output.append(total)
    return output


def face_interval(polynomial, lower, upper, coordinate, endpoint):
    # Substitution precedes aggregation, preserving cancellations on the face.
    reduced = {}
    for exponent, coefficient in polynomial.items():
        e = list(exponent); power = e[coordinate]; e[coordinate] = 0; e = tuple(e)
        reduced[e] = bounded(reduced.get(e, Q(0))+bounded(coefficient*bounded(endpoint**power)))
    low = high = Q(0)
    for exponent, coefficient in reduced.items():
        a = b = Q(1)
        for l, u, power in zip(lower, upper, exponent):
            a = bounded(a*bounded(l**power)); b = bounded(b*bounded(u**power))
        lo, hi = (coefficient*a, coefficient*b) if coefficient >= 0 else (coefficient*b, coefficient*a)
        low = bounded(low+bounded(lo)); high = bounded(high+bounded(hi))
    return low, high


def audit(data):
    names, complexes, rates, initial, vectors, box, coefficients, removed = parse(data)
    n = len(names); edges = sorted(rates)
    graph = graph_info(len(complexes), edges)
    declared = graph_info(len(complexes), {(r['source'],r['target']) for r in data['reactions']})
    rows = [[complexes[j][k]-complexes[i][k] for k in range(n)] for i,j in edges]
    rank, basis = nullspace(rows, n)
    polys = polynomials(complexes, rates, n)
    certificates = []; bounds = [None]*n
    for w in vectors:
        residuals = [dot(w, row) for row in rows]
        accepted = any(w) and all(x == 0 for x in residuals)
        total = dot(w, initial)
        proposed = [bounded(total/x) if accepted and x > 0 else None for x in w]
        if accepted:
            for i, value in enumerate(proposed):
                if value is not None: bounds[i] = min(value, bounds[i]) if bounds[i] is not None else value
        certificates.append(dict(weights=list(map(str,w)), accepted=accepted, residuals=list(map(str,residuals)),
                                 conserved_total=str(total) if accepted else None,
                                 coordinate_upper_bounds=[str(x) if x is not None else None for x in proposed]))
    box_report = dict(status='not_supplied', initial_contained=False, all_time_bounds_for_initial=False)
    if box:
        lower, upper = box; faces = []
        for i in range(n):
            for side, endpoint in [('lower',lower[i]), ('upper',upper[i])]:
                lo, hi = face_interval(polys[i], lower, upper, i, endpoint)
                passed = lo >= 0 if side == 'lower' else hi <= 0
                faces.append(dict(species=names[i], side=side, derivative_lower=str(lo), derivative_upper=str(hi), inward_verified=passed))
        accepted = all(f['inward_verified'] for f in faces)
        inside = all(l <= x <= u for l,x,u in zip(lower,initial,upper))
        box_report = dict(status='invariant_box_verified' if accepted else 'not_certified',
                          lower=list(map(str,lower)), upper=list(map(str,upper)), faces=faces,
                          initial_contained=inside, all_time_bounds_for_initial=accepted and inside,
                          scope='Exact polynomial face enclosures; no attraction, entry-time, equilibrium-convergence or source-polytope construction claim. Failure is inconclusive.')
    applicable = graph['weakly_reversible'] and all(x > 0 for x in initial)
    return dict(schema='reaction-network-audit-v1', model_bridge_to_external_simulator_verified=False,
                source_proof_verified=False, source_scope_matches_effective_model=applicable,
                source_numeric_epsilon=None, source_entry_time=None,
                source_scope='Older selected formal statement: initial-state-dependent all-time bounds. Newer manuscript: eventual classwise bounds with initial-dependent entry time. Neither proof executed.',
                normalized_coefficients=coefficients, omitted_no_effect_reactions=removed,
                effective_edges=[dict(source=i,target=j,coefficient=str(rates[i,j])) for i,j in edges],
                declared_graph=declared, effective_graph=graph,
                stoichiometric_rank=rank, conservation_basis=[list(map(str,v)) for v in basis],
                conserved_totals=[str(dot(v,initial)) for v in basis],
                deficiency=len(complexes)-len(graph['linkage_classes'])-rank,
                polynomial_vector_field=[[dict(exponent=list(e),coefficient=str(c)) for e,c in sorted(p.items())] for p in polys],
                initial_derivative=list(map(str,evaluate(polys, initial))),
                supplied_conservation_certificates=certificates,
                conservation_upper_bounds=[str(x) if x is not None else None for x in bounds],
                conservation_bounds_all_coordinates=all(x is not None for x in bounds),
                box=box_report,
                limits=dict(max_species=8,max_complexes=64,max_reactions=128,max_degree=8,input_rational_bits=32,intermediate_rational_bits=MAX_BITS,max_bytes=MAX_BYTES))


def load(raw):
    if len(raw) > MAX_BYTES: raise ValueError('input exceeds one MiB')
    def pairs(items):
        result = {}
        for key,value in items:
            if key in result: raise ValueError('duplicate JSON key')
            result[key] = value
        return result
    def invalid(value): raise ValueError('nonfinite JSON is invalid')
    return json.loads(raw, object_pairs_hook=pairs, parse_constant=invalid)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input',type=Path); parser.add_argument('--output',type=Path,required=True)
    args = parser.parse_args()
    if args.output.exists(): raise SystemExit('Refusing to replace existing output')
    with args.input.open('rb') as f: raw = f.read(MAX_BYTES+1)
    try: report = audit(load(raw))
    except ArithmeticError as exc:
        report = dict(schema='reaction-network-audit-v1',status='unknown_resource_limit',reason=str(exc),source_proof_verified=False)
    report.update(input_sha256=hashlib.sha256(raw).hexdigest(), checker_sha256=hashlib.sha256(Path(__file__).read_bytes()).hexdigest())
    args.output.mkdir(parents=True,exist_ok=False)
    with (args.output/'input.json').open('xb') as f:f.write(raw)
    with (args.output/'report.json').open('x') as f:json.dump(report,f,indent=2);f.write('\n')
    print(json.dumps({k:report[k] for k in ['status','source_scope_matches_effective_model','stoichiometric_rank','conservation_bounds_all_coordinates'] if k in report}))


if __name__ == '__main__': main()
