#!/usr/bin/env python3
"""Bounded exact factorization audit and a supplied-separator scalar scan.

This uses conventional finite-field arithmetic and the Rabin criterion, not
the new large-characteristic backend of source family 142. Coefficients are
canonical residues in ascending degree order. No third-party dependency.
"""
import argparse
import hashlib
import json
from pathlib import Path

MAX_PRIME = 2**31 - 1
MAX_DEGREE = 64
MAX_BYTES = 1024**2
MAX_WORK = 2_000_000


class Invalid(ValueError):
    pass


class Exhausted(Exception):
    pass


class Field:
    def __init__(self, p, work_limit=MAX_WORK):
        integer(p, 'p', 2, MAX_PRIME)
        integer(work_limit, 'work_limit', 1, MAX_WORK)
        self.p, self.work, self.limit = p, 0, work_limit
        if p != 2 and p % 2 == 0:
            raise Invalid('p is composite')
        d = 3
        while d*d <= p:
            self.charge()
            if p % d == 0:
                raise Invalid('p is composite')
            d += 2

    def charge(self, n=1):
        if self.work+n > self.limit:
            raise Exhausted('scalar-operation budget exhausted')
        self.work += n

    def sub(self, a, b):
        self.charge(max(len(a), len(b)))
        return trim([((a[i] if i < len(a) else 0) -
                      (b[i] if i < len(b) else 0)) % self.p
                     for i in range(max(len(a), len(b)))])

    def mul(self, a, b):
        self.charge(len(a)*len(b))
        c = [0]*(len(a)+len(b)-1)
        for i, x in enumerate(a):
            for j, y in enumerate(b):
                c[i+j] = (c[i+j]+x*y) % self.p
        return trim(c)

    def divmod(self, a, b):
        if b == [0]:
            raise Invalid('division by zero polynomial')
        r = a[:]
        q = [0]*max(1, len(a)-len(b)+1)
        inv = pow(b[-1], -1, self.p)
        while r != [0] and len(r) >= len(b):
            self.charge(len(b))
            k = len(r)-len(b)
            t = r[-1]*inv % self.p
            q[k] = t
            for i, x in enumerate(b):
                r[k+i] = (r[k+i]-t*x) % self.p
            r = trim(r)
        return trim(q), r

    def gcd(self, a, b):
        while b != [0]:
            a, b = b, self.divmod(a, b)[1]
        if a == [0]:
            return a
        self.charge(len(a))
        inv = pow(a[-1], -1, self.p)
        return [x*inv % self.p for x in a]

    def powmod(self, a, e, m):
        r = [1]
        a = self.divmod(a, m)[1]
        while e:
            if e & 1:
                r = self.divmod(self.mul(r, a), m)[1]
            e >>= 1
            if e:
                a = self.divmod(self.mul(a, a), m)[1]
        return r


def trim(a):
    while len(a) > 1 and a[-1] == 0:
        a.pop()
    return a


def integer(x, name, lo, hi):
    if type(x) is not int or not lo <= x <= hi:
        raise Invalid(f'{name} must be an integer from {lo} through {hi}')
    return x


def polynomial(x, p, name, nonzero=True):
    if type(x) is not list or not 1 <= len(x) <= MAX_DEGREE+1:
        raise Invalid(f'{name}: expected 1 through {MAX_DEGREE+1} coefficients')
    for a in x:
        integer(a, f'{name} coefficient', 0, p-1)
    if len(x) > 1 and x[-1] == 0:
        raise Invalid(f'{name}: leading zero is noncanonical')
    if nonzero and x == [0]:
        raise Invalid(f'{name}: zero polynomial is unsupported')
    return x[:]


def keys(obj, expected, name):
    if type(obj) is not dict or set(obj) != set(expected):
        raise Invalid(f'{name}: fields must be {sorted(expected)}')


def prime_divisors(n):
    result = []
    for d in range(2, n+1):
        if n % d == 0:
            result.append(d)
            while n % d == 0:
                n //= d
    return result


def irreducibility(f, field):
    """Rabin: x^(p^n)=x mod f and gcd(x^(p^(n/q))-x,f)=1."""
    n = len(f)-1
    if n < 1:
        raise Invalid('irreducibility needs a positive degree')
    divisors = prime_divisors(n)
    checkpoints = {n//q: q for q in divisors}
    x = field.divmod([0, 1], f)[1]
    h = x
    checks = []
    for i in range(1, n+1):
        h = field.powmod(h, field.p, f)
        if i in checkpoints:
            g = field.gcd(field.sub(h, x), f)
            checks.append(dict(degree_prime_divisor=checkpoints[i],
                               frobenius_iterations=i, gcd=g, coprime=g == [1]))
    residual = field.sub(h, x)
    return dict(irreducible=residual == [0] and all(c['coprime'] for c in checks),
                degree=n, frobenius_final_residual=residual, gcd_checks=checks)


def audit(doc, field):
    keys(doc, ['schema', 'p', 'polynomial', 'leading_coefficient', 'factors'], 'audit')
    f = polynomial(doc['polynomial'], field.p, 'polynomial')
    c = integer(doc['leading_coefficient'], 'leading_coefficient', 1, field.p-1)
    factors = doc['factors']
    if type(factors) is not list or len(factors) > MAX_DEGREE:
        raise Invalid('factors must be a list of at most 64 entries')
    seen, parsed, degree = set(), [], 0
    for item in factors:
        keys(item, ['polynomial', 'multiplicity'], 'factor')
        g = polynomial(item['polynomial'], field.p, 'factor polynomial')
        e = integer(item['multiplicity'], 'multiplicity', 1, MAX_DEGREE)
        if len(g) < 2 or g[-1] != 1:
            raise Invalid('factors must be monic with positive degree')
        if tuple(g) in seen:
            raise Invalid('duplicate factors must be merged by adding multiplicities')
        seen.add(tuple(g)); parsed.append((g, e)); degree += (len(g)-1)*e
    if degree > MAX_DEGREE:
        raise Invalid('expanded factor degree exceeds 64')
    product, checks = [c], []
    for g, e in parsed:
        checks.append(dict(polynomial=g, multiplicity=e, **irreducibility(g, field)))
        for _ in range(e):
            product = field.mul(product, g)
    matches = product == f
    lc_matches = c == f[-1]
    irreducible = all(r['irreducible'] for r in checks)
    accepted = matches and lc_matches and irreducible
    return dict(status='verified' if accepted else 'rejected',
                factorization_verified=accepted, multiply_back_matches=matches,
                leading_coefficient_matches=lc_matches, all_factors_irreducible=irreducible,
                reconstructed_polynomial=product, factor_checks=checks)


def scalar_scan(doc, field):
    keys(doc, ['schema', 'p', 'polynomial', 'separator', 'trial_limit'], 'scan')
    f = polynomial(doc['polynomial'], field.p, 'polynomial')
    b = polynomial(doc['separator'], field.p, 'separator', nonzero=False)
    limit = integer(doc['trial_limit'], 'trial_limit', 1, 100_000)
    if len(f) < 3 or f[-1] != 1:
        raise Invalid('scan needs a monic polynomial of degree at least two')
    derivative = trim([i*f[i] % field.p for i in range(1, len(f))])
    if field.gcd(f, derivative) != [1]:
        raise Invalid('scan needs a square-free polynomial')
    b = field.divmod(b, f)[1]
    if len(b) == 1:
        raise Invalid('separator must be nonscalar modulo the polynomial')
    if field.sub(field.powmod(b, field.p, f), b) != [0]:
        raise Invalid('separator must be Frobenius-fixed modulo the polynomial')
    for a in range(min(field.p, limit)):
        g = field.gcd(f, field.sub(b, [a]))
        if 1 < len(g) < len(f):
            q, r = field.divmod(f, g)
            if r != [0] or field.mul(q, g) != f:
                raise ArithmeticError('internal divisor witness failure')
            return dict(status='proper_divisor_verified', factorization_verified=False,
                        trials=a+1, scalar=a, divisor=g, cofactor=q, separator=b,
                        interpretation='One supplied separator, not a complete factorization algorithm')
    return dict(status='unknown_trial_limit', factorization_verified=False,
                trials=min(field.p, limit), separator=b,
                interpretation='No divisor found within this scan; not an irreducibility conclusion')


def evaluate(doc, work_limit=MAX_WORK):
    field = None
    try:
        if type(doc) is not dict:
            raise Invalid('input must be an object')
        if doc.get('schema') not in ('finite-field-factorization-audit-v1', 'finite-field-scalar-scan-v1'):
            raise Invalid('unknown schema')
        field = Field(doc.get('p'), work_limit)
        result = audit(doc, field) if doc['schema'].endswith('audit-v1') else scalar_scan(doc, field)
    except Invalid as exc:
        result = dict(status='invalid_input', factorization_verified=False, reason=str(exc))
    except Exhausted as exc:
        result = dict(status='unknown_resource_limit', factorization_verified=False, reason=str(exc))
    result.update(work_charged=field.work if field else None,
                  limits=dict(prime_max=MAX_PRIME, degree_max=MAX_DEGREE, work_limit=work_limit),
                  scope='Conventional exact finite-field checks; no source large-characteristic backend or proof acceptance')
    return result


def unique_object(pairs):
    result = {}
    for k, v in pairs:
        if k in result:
            raise Invalid('duplicate JSON key')
        result[k] = v
    return result


def decode(raw):
    if len(raw) > MAX_BYTES:
        raise Invalid('input exceeds one MiB')
    return json.loads(raw.decode('utf-8'), object_pairs_hook=unique_object,
                      parse_constant=lambda _: (_ for _ in ()).throw(Invalid('nonfinite JSON value')))


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', required=True, type=Path)
    parser.add_argument('--work-limit', type=int, default=MAX_WORK)
    args = parser.parse_args()
    if args.output.exists():
        raise FileExistsError('Refusing to overwrite an existing output')
    with args.input.open('rb') as stream:
        raw = stream.read(MAX_BYTES+1)
    try:
        result = evaluate(decode(raw), args.work_limit)
    except (Invalid, ValueError, UnicodeError, RecursionError) as exc:
        result = dict(status='invalid_input', factorization_verified=False, reason=str(exc))
    result.update(input_sha256=hashlib.sha256(raw).hexdigest(),
                  input_bytes_read=len(raw), input_truncated=len(raw) > MAX_BYTES,
                  checker_sha256=hashlib.sha256(Path(__file__).read_bytes()).hexdigest())
    with args.output.open('x') as stream:
        json.dump(result, stream, indent=2); stream.write('\n')
    print(json.dumps(dict(status=result['status'], output=str(args.output))))


if __name__ == '__main__':
    main()
