#!/usr/bin/env python3
"""Exact sparse Fourier reference for |u|^(p-1)u and grid aliasing.

Coefficients are physical Fourier coefficients, not the weighted Sobolev
coordinates of the Lean implementation. This does not integrate the NLS.
"""
import argparse
import hashlib
import json
from fractions import Fraction
from pathlib import Path

MAX_WORK = 200_000
MAX_SUPPORT = 4_000
MAX_BITS = 1_024
ZERO = (Fraction(0), Fraction(0))


def multiply(a, b):
    return (a[0]*b[0]-a[1]*b[1], a[0]*b[1]+a[1]*b[0])


def calculate(record, work_budget=MAX_WORK):
    if not isinstance(record, dict) or set(record) != {'dimension', 'band', 'grid', 'power', 'coefficients'}:
        raise ValueError('Require dimension, band, grid, power and coefficients only')
    d, k, n, p = (record[a] for a in ('dimension', 'band', 'grid', 'power'))
    if any(isinstance(a, bool) or not isinstance(a, int) for a in (d, k, n, p)):
        raise ValueError('Parameters must be integers, not booleans')
    if not (1 <= d <= 12 and 0 <= k <= 16 and 1 <= n <= 257 and n % 2 and n >= 2*k+1 and 1 <= p <= 15 and p % 2):
        raise ValueError('Unsupported dimension, band, odd grid or odd power')
    if isinstance(work_budget, bool) or not isinstance(work_budget, int) or not 1 <= work_budget <= MAX_WORK:
        raise ValueError('Invalid work budget')
    rows = record['coefficients']
    if not isinstance(rows, list) or len(rows) > 32:
        raise ValueError('At most 32 sparse input coefficients')
    f, seen = {}, set()
    for row in rows:
        if not isinstance(row, dict) or set(row) != {'frequency', 'real', 'imaginary'}:
            raise ValueError('Invalid coefficient record')
        freq = row['frequency']
        if not isinstance(freq, list) or len(freq) != d or any(isinstance(x, bool) or not isinstance(x, int) or abs(x) > k for x in freq):
            raise ValueError('Input frequency must be in the retained coordinate band')
        freq = tuple(freq)
        if freq in seen:
            raise ValueError('Duplicate input frequency')
        seen.add(freq)
        pair = []
        for field in ('real', 'imaginary'):
            raw = row[field]
            if not isinstance(raw, str) or len(raw) > 100:
                raise ValueError('Coefficient must be a bounded rational string')
            try:
                value = Fraction(raw)
            except (ValueError, ZeroDivisionError) as exc:
                raise ValueError('Invalid rational coefficient') from exc
            if max(abs(value.numerator).bit_length(), value.denominator.bit_length()) > 128:
                raise ValueError('Coefficient exceeds 128 input bits')
            pair.append(value)
        if tuple(pair) != ZERO:
            f[freq] = tuple(pair)
    work = 0
    result = dict(input_record_sha256=hashlib.sha256(json.dumps(record, sort_keys=True, separators=(',', ':')).encode()).hexdigest(),
                  source_connection=dict(family='371', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb', source_proof_verification='not_run'),
                  formal_program_verification='not_run', arithmetic='exact_complex_rational',
                  coefficient_convention='physical Fourier coefficients; no weighted Sobolev coordinates',
                  limits=dict(work_budget=work_budget, support_cap=MAX_SUPPORT, intermediate_rational_bits=MAX_BITS, wall_time_deadline='none'),
                  dimension=d, band=k, grid=n, power=p,
                  dense_grid_complex128_storage_bytes=n**d*16,
                  storage_interpretation='Illustration for a fully materialized tensor grid; not a runtime or memory lower bound for sparse/structured implementations.',
                  sufficient_central_projection_grid=n > (p+1)*k,
                  sufficient_full_reconstruction_grid=n > 2*p*k)

    def add(poly, freq, value):
        previous = poly.get(freq, ZERO)
        value = (previous[0]+value[0], previous[1]+value[1])
        if any(max(abs(v.numerator).bit_length(), v.denominator.bit_length()) > MAX_BITS for v in value):
            raise OverflowError('Intermediate rational bit cap exceeded')
        if value == ZERO:
            poly.pop(freq, None)
        else:
            poly[freq] = value
        if len(poly) > MAX_SUPPORT:
            raise OverflowError('Sparse support cap exceeded')

    def charge():
        nonlocal work
        if work >= work_budget:
            raise OverflowError('Coefficient-pair/folding work budget exhausted')
        work += 1

    def convolve(left, right):
        output = {}
        for a, av in left.items():
            for b, bv in right.items():
                charge()
                add(output, tuple(a[h]+b[h] for h in range(d)), multiply(av, bv))
        return output

    def encode(poly):
        return [dict(frequency=list(freq), real=str(value[0]), imaginary=str(value[1])) for freq, value in sorted(poly.items())]

    try:
        exact = dict(f)
        if p > 1:
            conjugate = {tuple(-x for x in freq):(value[0], -value[1]) for freq, value in f.items()}
            absolute_squared = convolve(f, conjugate)
            for _ in range((p-1)//2):
                exact = convolve(exact, absolute_squared)
        folded = {}
        for freq, value in exact.items():
            charge()
            representative = tuple((x+n//2) % n-n//2 for x in freq)
            add(folded, representative, value)
        central_exact = {a:v for a,v in exact.items() if all(abs(x) <= k for x in a)}
        central_folded = {a:v for a,v in folded.items() if all(abs(x) <= k for x in a)}
        difference = {}
        for a in central_exact.keys() | central_folded.keys():
            charge()
            v, w = central_folded.get(a, ZERO), central_exact.get(a, ZERO)
            add(difference, a, (v[0]-w[0], v[1]-w[1]))
    except OverflowError as exc:
        return dict(result, status='budget_exceeded', reason=str(exc), work_units=work,
                    central_projection_matches=None, full_polynomial_recovered=None,
                    interpretation='Comparison unknown; no partial convolution is certified.')
    return dict(result, status='complete', work_units=work,
                exact_nonlinearity=encode(exact), circular_grid_coefficients=encode(folded),
                exact_retained_projection=encode(central_exact), circular_retained_projection=encode(central_folded),
                circular_minus_exact_retained=encode(difference),
                central_projection_matches=not difference, full_polynomial_recovered=exact == folded,
                interpretation='Exact finite sparse polynomial and circular-fold comparison. This neither proves the source blowup theorem nor certifies a floating FFT, time integrator or chosen blowup power.')


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


if __name__ == '__main__':
    main()
