#!/usr/bin/env python3
"""Bounded exact first-column evaluation of family 116's rational hitting matrices.

This prototype accepts a visible division-free formula tree over Q. A nonzero
evaluation is a finite counterexample to a free-polynomial identity, independently
of universal hitting. Zero is never presented as independently proved identity.
"""
import argparse
from fractions import Fraction
import json
from pathlib import Path
import re

MODEL = dict(field='rational_characteristic_zero', algebra='free_associative', gates='division_free_tree')
MAX_GATES = 127
MAX_DIMENSION = 16384
MAX_WORK = 1000000
MAX_BITS = 8192


class Exhausted(Exception):
    pass


class Meter:
    def __init__(self, limit, bits):
        self.limit, self.bits, self.used = limit, bits, 0

    def charge(self):
        if self.used >= self.limit:
            raise Exhausted('arithmetic budget exhausted')
        self.used += 1

    def checked(self, value):
        if max(abs(value.numerator).bit_length(), value.denominator.bit_length()) > self.bits:
            raise Exhausted('intermediate rational bit cap exceeded')
        return value

    def add(self, a, b):
        self.charge()
        return self.checked(a+b)

    def mul(self, a, b):
        self.charge()
        return self.checked(a*b)

    def div(self, a, b):
        self.charge()
        return self.checked(a/b)


def integer(value, name, low, high):
    if type(value) is not int or not low <= value <= high:
        raise ValueError(f'{name} must be an integer in [{low},{high}]')
    return value


def rational(text):
    if not isinstance(text, str) or len(text) > 180 or not re.fullmatch(r'-?(0|[1-9][0-9]*)(/[1-9][0-9]*)?', text):
        raise ValueError('Scalar must be a canonical rational string')
    result = Fraction(text)
    if str(result) != text or max(abs(result.numerator).bit_length(), result.denominator.bit_length()) > 256:
        raise ValueError('Scalar is not canonical or exceeds 256-bit input cap')
    return result


def parse_formula(node, variables, count, depth=0):
    count[0] += 1
    if count[0] > MAX_GATES or depth > 32:
        raise ValueError('Formula exceeds gate/depth cap')
    if not isinstance(node, dict) or len(node) != 1:
        raise ValueError('Each formula node must have exactly one gate')
    gate, value = next(iter(node.items()))
    if gate == 'scalar':
        return (gate, rational(value))
    if gate == 'var':
        return (gate, integer(value, 'variable index', 1, variables))
    if gate not in ('add', 'mul') or not isinstance(value, list) or len(value) != 2:
        raise ValueError('Only scalar, var, binary add and ordered mul gates are supported')
    return (gate, parse_formula(value[0], variables, count, depth+1),
            parse_formula(value[1], variables, count, depth+1))


def operator(vector, i, meter, exceptional=False):
    if exceptional:
        return list(vector)  # The source's n=s=1 tuple is [1], not the nilpotent truncation.
    result = [Fraction(0)] * len(vector)
    quotient = Fraction(0)
    for r in range(1, len(vector)):
        # (z+i)c(z)=v(z), followed by integration with zero constant term.
        quotient = meter.div(meter.add(vector[r-1], -quotient), i)
        result[r] = meter.div(quotient, r)
    return result


def apply_formula(formula, vector, meter, exceptional):
    gate = formula[0]
    if gate == 'scalar':
        return [meter.mul(formula[1], v) for v in vector]
    if gate == 'var':
        return operator(vector, formula[1], meter, exceptional)
    if gate == 'add':
        left = apply_formula(formula[1], vector, meter, exceptional)
        right = apply_formula(formula[2], vector, meter, exceptional)
        return [meter.add(a,b) for a,b in zip(left,right)]
    right = apply_formula(formula[2], vector, meter, exceptional)
    return apply_formula(formula[1], right, meter, exceptional)


def audit(record, coefficients=None, work_limit=MAX_WORK, bit_limit=MAX_BITS):
    if not isinstance(record, dict) or set(record) != {'variables', 'formula', 'model'} or record['model'] != MODEL:
        raise ValueError('Require variables, formula and the exact supported model declaration')
    variables = integer(record['variables'], 'variables', 1, 8)
    work_limit = integer(work_limit, 'work limit', 0, MAX_WORK)
    bit_limit = integer(bit_limit, 'intermediate bit cap', 1, MAX_BITS)
    count = [0]
    formula = parse_formula(record['formula'], variables, count)
    gates = count[0]
    exceptional = variables == gates == 1
    source_dimension = 1 if exceptional else 2*variables*gates*gates
    dimension = source_dimension if coefficients is None else min(source_dimension, integer(coefficients, 'coefficient limit', 1, MAX_DIMENSION))
    base = dict(variables=variables, gates=gates, source_dimension=source_dimension,
                requested_dimension=dimension, completed_dimension=None, requested_full_source_dimension=(dimension==source_dimension),
                model=MODEL, independent_identity_conclusion='unknown',
                source_universal_hitting_acceptance='not_independently_verified',
                implementation='Visible formula tree; structured rational matrix action, not an opaque black-box oracle',
                budget_scope='Counts selected rational add/multiply/divide calls only; excludes parsing, negation, bit arithmetic, allocation and wall time')
    if dimension > MAX_DIMENSION:
        return dict(base, status='unknown', reason='dimension cap exceeded', arithmetic_charges=0)
    meter = Meter(work_limit, bit_limit)
    vector = [Fraction(0)]*dimension
    vector[0] = Fraction(1)
    try:
        result = apply_formula(formula, vector, meter, exceptional)
    except Exhausted as exc:
        return dict(base, status='unknown', reason=str(exc), arithmetic_charges=meter.used)
    nonzero = next(((i,v) for i,v in enumerate(result) if v), None)
    if nonzero:
        row, value = nonzero
        return dict(base, status='nonzero_polynomial_witness', completed_dimension=dimension, independent_identity_conclusion='not_an_identity_in_the_declared_free_rational_algebra',
                    witness=dict(row=row, column=0, value=str(value), dimension=dimension,
                                 generator='[1] for n=s=1; otherwise entry (-1)^(r-q-1)/(r*i^(r-q)) for q<r, zero otherwise'),
                    arithmetic_charges=meter.used)
    return dict(base, completed_dimension=dimension, status='zero_at_source_tuple_conditional' if dimension==source_dimension else 'zero_at_truncation_unknown',
                conditional_interpretation='A full-dimension zero first column implies identity only under the source paper detection argument and correct implementation; neither is formally accepted here',
                arithmetic_charges=meter.used)


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--coefficients', type=int)
    parser.add_argument('--work-limit', type=int, default=MAX_WORK)
    parser.add_argument('--bit-limit', type=int, default=MAX_BITS)
    args=parser.parse_args()
    try:
        if args.input.stat().st_size > 2*1024*1024:
            raise ValueError('Input exceeds 2 MiB cap')
        record=json.loads(args.input.read_text())
        result=audit(record,args.coefficients,args.work_limit,args.bit_limit)
    except (ValueError,RecursionError) as exc:
        print(json.dumps(dict(status='invalid_input',reason=str(exc))))
        raise SystemExit(2)
    print(json.dumps(result,indent=2))


if __name__=='__main__': main()
