#!/usr/bin/env python3
"""Small exact natural-unit-interval chromatic elementary-basis reference.

Conventional finite coloring enumeration and rational basis conversion,
not the source packet procedure or an all-r formal witness certificate.
"""
import argparse
import hashlib
import itertools
import json
from fractions import Fraction
from pathlib import Path

MAX_WORK = 2_000_000


def partitions(n, ceiling=None):
    if n == 0:
        yield ()
    else:
        for first in range(min(n, ceiling if ceiling is not None else n), 0, -1):
            for tail in partitions(n-first,first):
                yield (first,)+tail


def compositions(total, width):
    if width == 1:
        yield (total,)
    else:
        for first in range(total+1):
            for tail in compositions(total-first,width-1):
                yield (first,)+tail


def calculate(record, work_budget=MAX_WORK):
    if not isinstance(record,dict) or not {'h'} <= set(record) <= {'h','witness'}:
        raise ValueError('Provide h and optional witness only')
    h=record['h']
    if not isinstance(h,list) or not 1 <= len(h) <= 6:
        raise ValueError('Require one to six vertices')
    n=len(h)
    if any(isinstance(v,bool) or not isinstance(v,int) or not i <= v < n for i,v in enumerate(h)) or any(h[i]>h[i+1] for i in range(n-1)):
        raise ValueError('h must be extensive, nondecreasing and zero-indexed')
    if isinstance(work_budget,bool) or not isinstance(work_budget,int) or not 1 <= work_budget <= MAX_WORK:
        raise ValueError('Invalid work budget')
    edges=[(i,j) for i in range(n) for j in range(i+1,h[i]+1)]
    edge_set=set(edges)
    parts=list(partitions(n))
    work=0
    result=dict(input_record_sha256=hashlib.sha256(json.dumps(record,sort_keys=True,separators=(',',':')).encode()).hexdigest(),
                source_connection=dict(family='169',source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',source_proof_verification='not_run'),
                formal_program_verification='not_run',arithmetic='exact_integer_and_rational',vertices=n,edges=[list(e) for e in edges],colors=n,
                limits=dict(vertices=6,work_budget=work_budget,wall_time_deadline='none'),
                work_unit_definition='Color assignments/edge probes, polynomial product terms, coefficient comparisons and rational elimination operations; not bit operations or elapsed time.')

    def charge():
        nonlocal work
        if work>=work_budget:
            raise OverflowError('Exact enumeration/conversion budget exhausted')
        work+=1

    def elementary(part):
        poly={(0,)*n:1}
        for size in part:
            following={}
            for selected in itertools.combinations(range(n),size):
                selected=set(selected)
                for exp,value in poly.items():
                    charge()
                    key=tuple(exp[i]+(i in selected) for i in range(n))
                    following[key]=following.get(key,0)+value
            poly=following
        return poly

    def adjacent_edge(a,b):
        return (min(a,b),max(a,b)) in edge_set

    # Parse and check user-supplied permutation labels before bounded work.
    supplied=None
    if 'witness' in record:
        rows=record['witness']
        if not isinstance(rows,list) or len(rows)>720:
            raise ValueError('Invalid witness list')
        supplied={}
        for row in rows:
            if not isinstance(row,dict) or set(row)!={'permutation','partition'}:
                raise ValueError('Invalid witness row')
            sigma,part=row['permutation'],row['partition']
            if not isinstance(sigma,list) or any(isinstance(x,bool) or not isinstance(x,int) for x in sigma) or sorted(sigma)!=list(range(n)):
                raise ValueError('Invalid witness permutation')
            if not isinstance(part,list) or not part or any(isinstance(x,bool) or not isinstance(x,int) or x<=0 for x in part) or sum(part)!=n or part!=sorted(part,reverse=True):
                raise ValueError('Invalid partition')
            sigma=tuple(sigma)
            if sigma in supplied or any(sigma[i+1]<sigma[i] and not adjacent_edge(sigma[i],sigma[i+1]) for i in range(n-1)):
                raise ValueError('Duplicate or noneligible witness permutation')
            supplied[sigma]=tuple(part)
    try:
        colored={}
        proper_count=0
        for color in itertools.product(range(n),repeat=n):
            charge()
            proper=True
            degree=0
            for i,j in edges:
                charge()
                if color[i]==color[j]:
                    proper=False
                    break
                degree+=color[i]<color[j]
            if proper:
                proper_count+=1
                exp=tuple(color.count(c) for c in range(n))
                colored[(exp,degree)]=colored.get((exp,degree),0)+1
        degree_count=len(edges)+1
        monomial={part:[colored.get((part+(0,)*(n-len(part)),q),0) for q in range(degree_count)] for part in parts}
        symmetric=True
        for exp in compositions(n,n):
            part=tuple(sorted((v for v in exp if v),reverse=True))
            for q in range(degree_count):
                charge()
                symmetric &= colored.get((exp,q),0)==monomial[part][q]
        if not symmetric:
            return dict(result,status='model_mismatch',work_units=work,elementary_coefficients=None,reason='Complete finite polynomial is not symmetric under color-variable permutation')
        basis={part:elementary(part) for part in parts}
        width=len(parts)
        matrix=[[Fraction(basis[col].get(row+(0,)*(n-len(row)),0)) for col in parts]+[Fraction(v) for v in monomial[row]] for row in parts]
        for col in range(width):
            pivot=next((i for i in range(col,width) if matrix[i][col]),None)
            if pivot is None:
                raise ArithmeticError('Elementary basis matrix unexpectedly singular')
            matrix[col],matrix[pivot]=matrix[pivot],matrix[col]
            divisor=matrix[col][col]
            for j in range(col,width+degree_count):
                charge(); matrix[col][j]/=divisor
            for i in range(width):
                if i==col:
                    continue
                factor=matrix[i][col]
                for j in range(col,width+degree_count):
                    charge(); matrix[i][j]-=factor*matrix[col][j]
        elementary_coeff={part:matrix[i][width:] for i,part in enumerate(parts)}
        reconstructed={}
        for part,coeffs in elementary_coeff.items():
            for q,value in enumerate(coeffs):
                if not value:
                    continue
                for exp,mult in basis[part].items():
                    charge()
                    key=(exp,q)
                    reconstructed[key]=reconstructed.get(key,Fraction(0))+value*mult
        reconstructed={key:value for key,value in reconstructed.items() if value}
        reconstruction_matches=reconstructed==colored
        if not reconstruction_matches:
            raise ArithmeticError('Complete polynomial reconstruction failed')
        eligible={}
        for sigma in itertools.permutations(range(n)):
            charge()
            if any(sigma[i+1]<sigma[i] and not adjacent_edge(sigma[i],sigma[i+1]) for i in range(n-1)):
                continue
            degree=sum(sigma[j]<sigma[i] and adjacent_edge(sigma[i],sigma[j]) for i in range(n) for j in range(i+1,n))
            eligible[sigma]=degree
        witness_status='not_supplied'
        differences=[]
        if supplied is not None:
            if set(supplied)!=set(eligible):
                witness_status='incomplete'
            else:
                totals={part:[0]*degree_count for part in parts}
                for sigma,part in supplied.items():
                    totals[part][eligible[sigma]]+=1
                for part in parts:
                    for q in range(degree_count):
                        charge()
                        if totals[part][q]!=elementary_coeff[part][q]:
                            differences.append(dict(partition=list(part),q_degree=q,supplied_count=totals[part][q],expected_coefficient=str(elementary_coeff[part][q])))
                witness_status='finite_polynomial_matches' if not differences else 'finite_polynomial_mismatch'
    except OverflowError as exc:
        return dict(result,status='budget_exceeded',reason=str(exc),work_units=work,elementary_coefficients=None,witness_status='unknown',interpretation='Complete expansion unknown; no partial coefficient or positivity certificate.')
    return dict(result,status='complete',work_units=work,proper_colorings_with_n_colors=proper_count,
                elementary_coefficients=[dict(partition=list(part),q_coefficients=[str(v) for v in coeffs]) for part,coeffs in elementary_coeff.items() if any(coeffs)],
                monomial_coefficients=[dict(partition=list(part),q_coefficients=coeffs) for part,coeffs in monomial.items() if any(coeffs)],
                coefficients_nonnegative_integers=all(v.denominator==1 and v>=0 for coeffs in elementary_coeff.values() for v in coeffs),
                complete_finite_polynomial_reconstructed=True,eligible_permutations=len(eligible),witness_status=witness_status,witness_differences=differences,
                interpretation='Exact degree-n polynomial in n colors and conventional basis conversion. Supplied witness status checks this finite polynomial only; no source packet algorithm, all-r formal proof, uniform constructive witness, geometric basis or polynomial runtime is certified.')


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()
