#!/usr/bin/env python3
"""Bounded exact finite-matrix reference; no source proof/program acceptance.

Complex entries are pairs of canonical rational strings, e.g. ["1/2", "0"].
Coefficients[k] is a square matrix B_k. Evaluation is sum A^k tensor B_k.
Exact PSD threshold tests are independent of the Crouzeix source claim.
Support halfspaces and a covered-cell Lipschitz bound use rational arithmetic.
"""
import argparse
import hashlib
import json
from fractions import Fraction as F
from math import isqrt

MODEL = {"matrix_domain": "finite_complex_rational", "norm": "Euclidean_operator_2",
         "evaluation": "sum_A_power_k_tensor_B_k_base_first",
         "polynomial_variable": "one_complex_scalar", "coefficient_order": "ascending"}
ZERO, ONE = (F(0), F(0)), (F(1), F(0))
PRIOR_CONSTANT = F(483, 200)  # 2.415 > 1+sqrt(2), since (283/200)^2 > 2.


class Limit(Exception):
    pass


class Budget:
    def __init__(self, maximum=2_000_000, bits=4096):
        self.maximum, self.bits, self.used = maximum, bits, 0

    def charge(self, *values):
        self.used += 1
        if self.used > self.maximum:
            raise Limit("charged_arithmetic_budget_exhausted")
        for v in values:
            if max(v.numerator.bit_length(), v.denominator.bit_length()) > self.bits:
                raise Limit("intermediate_rational_bit_limit")

    def add(self, x, y):
        z = x + y
        self.charge(z)
        return z

    def mul(self, x, y):
        z = x * y
        self.charge(z)
        return z

    def div(self, x, y):
        z = x / y
        self.charge(z)
        return z


def cadd(x, y, b):
    return b.add(x[0], y[0]), b.add(x[1], y[1])


def cmul(x, y, b):
    return (b.add(b.mul(x[0], y[0]), -b.mul(x[1], y[1])),
            b.add(b.mul(x[0], y[1]), b.mul(x[1], y[0])))


def conjugate(z):
    return z[0], -z[1]


def eye(n):
    return [[ONE if i == j else ZERO for j in range(n)] for i in range(n)]


def adjoint(a):
    return [[conjugate(a[j][i]) for j in range(len(a))] for i in range(len(a[0]))]


def matrix_mul(a, c, b):
    out = []
    for row in a:
        output = []
        for j in range(len(c[0])):
            s = ZERO
            for k, x in enumerate(row):
                s = cadd(s, cmul(x, c[k][j], b), b)
            output.append(s)
        out.append(output)
    return out


def realify(h):
    n = len(h)
    return [[h[i][j][0] for j in range(n)] + [-h[i][j][1] for j in range(n)] for i in range(n)] + [
        [h[i][j][1] for j in range(n)] + [h[i][j][0] for j in range(n)] for i in range(n)]


def psd(a, b):
    """Exact symmetric PSD by Schur complements, including zero pivots.

    No positive-pivot threshold or floating arithmetic is used. A zero diagonal
    in a PSD matrix requires its entire row to be zero.
    """
    n = len(a)
    if any(len(row) != n for row in a) or any(a[i][j] != a[j][i] for i in range(n) for j in range(n)):
        raise ValueError("PSD input must be real symmetric square")
    s = [list(row) for row in a]
    pivots = []
    for k in range(n):
        pivot = s[k][k]
        b.charge(pivot)
        pivots.append(str(pivot))
        if pivot < 0:
            return {"psd": False, "pivots": pivots, "failure": "negative_pivot", "index": k}
        if pivot == 0:
            if any(s[k][j] != 0 for j in range(k + 1, n)):
                return {"psd": False, "pivots": pivots, "failure": "nonzero_zero_pivot_row", "index": k}
            continue
        for i in range(k + 1, n):
            for j in range(i, n):
                t = b.div(b.mul(s[i][k], s[k][j]), pivot)
                s[i][j] = b.add(s[i][j], -t)
                s[j][i] = s[i][j]
    return {"psd": True, "pivots": pivots, "failure": None}


def sqrt_upper(q, b):
    if q < 0:
        raise ValueError("negative square-root argument")
    scale = 1 << 24
    num = q.numerator * scale * scale
    den = q.denominator
    k = isqrt(num // den)
    if k * k * den < num:
        k += 1
    value = F(k, scale)
    b.charge(q, value)
    return value


def frobenius_upper(a, b):
    s = F(0)
    for row in a:
        for x, y in row:
            s = b.add(s, b.add(b.mul(x, x), b.mul(y, y)))
    return sqrt_upper(s, b)


def polynomial_at(coefficients, z, b):
    result = [list(row) for row in coefficients[-1]]
    for coefficient in reversed(coefficients[:-1]):
        result = [[cadd(cmul(z, result[i][j], b), coefficient[i][j], b)
                   for j in range(len(result))] for i in range(len(result))]
    return result


def tensor_evaluation(a, coefficients, b):
    n, m = len(a), len(coefficients[0])
    out = [[ZERO for _ in range(n * m)] for _ in range(n * m)]
    power = eye(n)
    for k, coefficient in enumerate(coefficients):
        for i in range(n):
            for j in range(n):
                for r in range(m):
                    for s in range(m):
                        out[i*m+r][j*m+s] = cadd(out[i*m+r][j*m+s], cmul(power[i][j], coefficient[r][s], b), b)
        if k + 1 < len(coefficients):
            power = matrix_mul(power, a, b)
    return out


def wire_matrix(a):
    return [[[str(x), str(y)] for x, y in row] for row in a]


def direct_threshold(a, coefficients, tolerance, b):
    evaluation = tensor_evaluation(a, coefficients, b)
    gram = matrix_mul(adjoint(evaluation), evaluation, b)
    t2 = b.mul(tolerance, tolerance)
    difference = [[(b.add(t2 if i == j else F(0), -gram[i][j][0]), -gram[i][j][1])
                   for j in range(len(gram))] for i in range(len(gram))]
    certificate = psd(realify(difference), b)
    digest = hashlib.sha256(json.dumps(wire_matrix(difference), separators=(",", ":")).encode()).hexdigest()
    return {"status": "complete", "norm_at_most_tolerance": certificate["psd"],
            "criterion": "tolerance_squared_I_minus_E_adjoint_E_is_PSD",
            "difference_matrix_sha256": digest, "realification_certificate": certificate,
            "formal_program_verification": "not_run", "source_Crouzeix_theorem_required": False,
            "evaluation_dimension": len(evaluation)}


def directions():
    result = [(F(1), F(0)), (F(-1), F(0)), (F(0), F(1)), (F(0), F(-1))]
    for x, y, denominator in ((3, 4, 5), (4, 3, 5), (5, 12, 13), (12, 5, 13)):
        result += [(F(sx*x, denominator), F(sy*y, denominator)) for sx in (-1, 1) for sy in (-1, 1)]
    return result


def support_matrix(a, direction, b):
    n = len(a)
    result = []
    for i in range(n):
        row = []
        for j in range(n):
            value = cadd(cmul(conjugate(direction), a[i][j], b), cmul(direction, conjugate(a[j][i]), b), b)
            row.append((b.div(value[0], F(2)), b.div(value[1], F(2))))
        result.append(row)
    return result


def support_certificate(a, direction, b, bisections):
    h = support_matrix(a, direction, b)
    n = len(h)
    low = max(h[i][i][0] for i in range(n))
    high = max(h[i][i][0] + sum(abs(h[i][j][0]) + abs(h[i][j][1]) for j in range(n) if j != i) for i in range(n))

    def check(value):
        shifted = [[(b.add(value if i == j else F(0), -h[i][j][0]), -h[i][j][1])
                    for j in range(n)] for i in range(n)]
        return psd(realify(shifted), b)

    certificate = check(high)
    if not certificate["psd"]:
        raise ArithmeticError("initial Gershgorin support upper bound failed")
    for _ in range(bisections):
        if high == low:
            break
        mid = b.div(b.add(high, low), F(2))
        tested = check(mid)
        if tested["psd"]:
            high, certificate = mid, tested
        else:
            low = mid
    return {"direction": [str(v) for v in direction], "upper": str(high),
            "criterion": "upper_I_minus_Re_conjugate_direction_A_is_PSD",
            "realification_certificate": certificate}


def enclosure_bound(a, coefficients, b, grid, bisections):
    certificates = [support_certificate(a, direction, b, bisections) for direction in directions()]
    planes = [(F(c["direction"][0]), F(c["direction"][1]), F(c["upper"])) for c in certificates]
    xlo, xhi, ylo, yhi = -planes[1][2], planes[0][2], -planes[3][2], planes[2][2]
    if xlo > xhi or ylo > yhi:
        raise ArithmeticError("invalid numerical-range enclosure")
    radius = sqrt_upper(b.add(b.mul(max(abs(xlo), abs(xhi)), max(abs(xlo), abs(xhi))),
                             b.mul(max(abs(ylo), abs(yhi)), max(abs(ylo), abs(yhi)))), b)
    norms = [frobenius_upper(c, b) for c in coefficients]
    lipschitz, rpower = F(0), F(1)
    for k in range(1, len(coefficients)):
        lipschitz = b.add(lipschitz, b.mul(F(k), b.mul(norms[k], rpower)))
        rpower = b.mul(rpower, radius)
    nx, ny = (1 if xhi == xlo else grid), (1 if yhi == ylo else grid)
    dx, dy = b.div(xhi-xlo, F(nx)), b.div(yhi-ylo, F(ny))
    halfx, halfy = b.div(dx, F(2)), b.div(dy, F(2))
    cell_radius = sqrt_upper(b.add(b.mul(halfx, halfx), b.mul(halfy, halfy)), b)
    correction = b.mul(lipschitz, cell_radius)
    maximum, retained, rejected = F(0), 0, 0
    for i in range(nx):
        left, right = b.add(xlo, b.mul(F(i), dx)), b.add(xlo, b.mul(F(i+1), dx))
        for j in range(ny):
            bottom, top = b.add(ylo, b.mul(F(j), dy)), b.add(ylo, b.mul(F(j+1), dy))
            outside = False
            for x, y, limit in planes:
                minimum = b.add(b.mul(x, left if x >= 0 else right), b.mul(y, bottom if y >= 0 else top))
                if minimum > limit:
                    outside = True
                    break
            if outside:
                rejected += 1
                continue
            retained += 1
            center = (b.add(left, halfx), b.add(bottom, halfy))
            upper = b.add(frobenius_upper(polynomial_at(coefficients, center, b), b), correction)
            maximum = max(maximum, upper)
    if retained == 0:
        raise ArithmeticError("no covered cells for nonempty numerical range")
    sharp, prior = b.mul(F(2), maximum), b.mul(PRIOR_CONSTANT, maximum)
    return {"status": "complete", "support_halfspaces": certificates,
            "rectangle": [str(v) for v in (xlo, xhi, ylo, yhi)], "grid_shape": [nx, ny],
            "retained_cells": retained, "rejected_cells": rejected,
            "global_radius_upper": str(radius), "coefficient_Frobenius_uppers": [str(v) for v in norms],
            "Lipschitz_upper": str(lipschitz), "cell_radius_upper": str(cell_radius),
            "polynomial_supremum_upper": str(maximum),
            "coverage_method": "closed rectangle cells; reject only if a halfspace excludes the entire cell; Frobenius center upper plus Lipschitz times radius",
            "source_constant_two_upper": str(sharp), "source_constant_two_stage": "conditional_on_unverified_repository_theorem",
            "published_prior_rational_constant": str(PRIOR_CONSTANT), "published_prior_upper": str(prior),
            "published_prior_reference": "https://epubs.siam.org/doi/10.1137/17M1116672",
            "published_prior_proof_execution_here": "not_run"}


def parse_rational(value, positive=False):
    if not isinstance(value, str) or len(value) > 90:
        raise ValueError("rational must be a bounded canonical string")
    try:
        parsed = F(value)
    except (ValueError, ZeroDivisionError) as error:
        raise ValueError("invalid rational") from error
    if str(parsed) != value or max(parsed.numerator.bit_length(), parsed.denominator.bit_length()) > 128:
        raise ValueError("noncanonical or over-128-bit rational")
    if positive and parsed <= 0:
        raise ValueError("tolerance must be positive")
    return parsed


def parse_matrix(value, maximum):
    if not isinstance(value, list) or not 1 <= len(value) <= maximum:
        raise ValueError("matrix dimension outside bounded reference")
    n = len(value)
    result = []
    for row in value:
        if not isinstance(row, list) or len(row) != n:
            raise ValueError("matrix must be square")
        output = []
        for z in row:
            if not isinstance(z, list) or len(z) != 2:
                raise ValueError("complex entry must have two rational strings")
            output.append(tuple(parse_rational(v) for v in z))
        result.append(output)
    return result


def run(data, grid=16, bisections=12, work=2_000_000, bits=4096):
    if not isinstance(data, dict) or set(data) != {"matrix", "coefficients", "tolerance", "model"} or data["model"] != MODEL:
        raise ValueError("exact declared finite model and four input fields required")
    if type(grid) is not int or not 1 <= grid <= 32 or type(bisections) is not int or not 0 <= bisections <= 16:
        raise ValueError("grid 1..32 and support bisections 0..16 required")
    if type(work) is not int or not 1 <= work <= 2_000_000 or type(bits) is not int or not 128 <= bits <= 4096:
        raise ValueError("work 1..2m and intermediate bits 128..4096 required")
    a = parse_matrix(data["matrix"], 8)
    if not isinstance(data["coefficients"], list) or not 1 <= len(data["coefficients"]) <= 7:
        raise ValueError("one through seven ascending matrix coefficients required")
    coefficients = [parse_matrix(c, 4) for c in data["coefficients"]]
    if any(len(c) != len(coefficients[0]) for c in coefficients):
        raise ValueError("coefficient dimensions must agree")
    tolerance = parse_rational(data["tolerance"], positive=True)
    b = Budget(work, bits)
    report = {"model": MODEL, "base_dimension": len(a), "coefficient_dimension": len(coefficients[0]),
              "degree": len(coefficients)-1, "tolerance": str(tolerance),
              "direct_threshold": {"status": "unknown"}, "numerical_range_bound": {"status": "unknown"},
              "limitations": "Finite exact rational input only; input-to-physical-model bridge, formal program/source proof, unrestricted runtime/memory and customer value unverified. Charged arithmetic excludes parsing, integer bit/gcd cost, native allocation, support setup and serialization; no wall-time bound."}
    try:
        report["direct_threshold"] = direct_threshold(a, coefficients, tolerance, b)
        report["direct_threshold"]["charged_work_at_completion"] = b.used
        bound = enclosure_bound(a, coefficients, b, grid, bisections)
        bound["source_conditional_upper_meets_tolerance"] = F(bound["source_constant_two_upper"]) <= tolerance
        bound["published_prior_upper_meets_tolerance"] = F(bound["published_prior_upper"]) <= tolerance
        bound["upper_exceeds_tolerance_interpretation"] = "inconclusive about actual norm; direct finite threshold is a separate exact result"
        report["numerical_range_bound"] = bound
        report["status"] = "complete"
    except Limit as error:
        report["status"] = "budget_exhausted"
        report["exhaustion_reason"] = str(error)
    report["charged_arithmetic"] = b.used
    report["configured_budget"] = work
    return report


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("input")
    parser.add_argument("--grid", type=int, default=16)
    parser.add_argument("--bisections", type=int, default=12)
    parser.add_argument("--work", type=int, default=2_000_000)
    parser.add_argument("--bits", type=int, default=4096)
    args = parser.parse_args()
    with open(args.input) as stream:
        data = json.load(stream)
    print(json.dumps(run(data, args.grid, args.bisections, args.work, args.bits), indent=2))


if __name__ == "__main__":
    main()
