#!/usr/bin/env python3
"""Shared-switch point/inverse replay and tiny exact coordinate-sweep laws.

The replay accepts explicit, possibly sparse shared bit assignments. It does
not generate independent randomness, implement a PRG, or grant a mixing bound.
"""
import argparse
import hashlib
import itertools
import json
from collections import defaultdict
from fractions import Fraction
from math import factorial
from pathlib import Path

MAX_WORK=4_000_000
MAX_STATES=50_000
MAX_TRACE=5_000


def integer(value, lower, upper, name):
    if isinstance(value,bool) or not isinstance(value,int) or not lower<=value<=upper:
        raise ValueError('Invalid '+name)
    return value


def pair_id(position, bit_position):
    return ((position>>(bit_position+1))<<bit_position)|(position&((1<<bit_position)-1))


def induced(permutation, target_size):
    """Delete outside-domain elements from the cycles of a full permutation."""
    output=[]
    for start in range(target_size):
        current=permutation[start]
        while current>=target_size:
            current=permutation[current]
        output.append(current)
    return tuple(output)


def replay(record, work_budget):
    if set(record)!={'mode','dimension','sweeps','target_size','queries','bits'}:
        raise ValueError('Invalid replay fields')
    d=integer(record['dimension'],1,63,'dimension')
    sweeps=integer(record['sweeps'],0,64,'sweeps')
    n=1<<d
    m=integer(record['target_size'],1,n,'target_size')
    queries=record['queries']
    if not isinstance(queries,list) or len(queries)>256:
        raise ValueError('At most 256 point/inverse queries')
    for q in queries:
        if not isinstance(q,dict) or set(q)!={'index','direction'} or q['direction'] not in ('forward','inverse'):
            raise ValueError('Invalid query')
        integer(q['index'],0,m-1,'query index')
    rows=record['bits']
    if not isinstance(rows,list) or len(rows)>100_000:
        raise ValueError('At most 100000 explicit shared bits')
    bits={}
    for row in rows:
        if not isinstance(row,dict) or set(row)!={'sweep','coordinate','pair','bit'}:
            raise ValueError('Invalid bit row')
        key=(integer(row['sweep'],0,sweeps-1,'bit sweep'),integer(row['coordinate'],0,d-1,'bit coordinate'),integer(row['pair'],0,n//2-1,'bit pair'))
        value=integer(row['bit'],0,1,'switch bit')
        if key in bits:
            raise ValueError('Duplicate shared bit key')
        bits[key]=value
    work=0
    trace=[]
    trace_omitted=0
    accessed=set()
    outputs=[]

    def apply_once(position, direction, query_index, hop):
        nonlocal work,trace_omitted
        sweep_order=range(sweeps) if direction=='forward' else range(sweeps-1,-1,-1)
        coordinate_order=range(d) if direction=='forward' else range(d-1,-1,-1)
        for sweep in sweep_order:
            for coordinate in coordinate_order:
                if work>=work_budget:
                    raise OverflowError('Shared switch lookup budget exhausted')
                work+=1
                bit_position=d-1-coordinate
                key=(sweep,coordinate,pair_id(position,bit_position))
                if key not in bits:
                    raise KeyError(key)
                accessed.add(key)
                following=position^((1<<bit_position) if bits[key] else 0)
                if len(trace)<MAX_TRACE:
                    trace.append(dict(query=query_index,cycle_hop=hop,sweep=sweep,coordinate=coordinate,pair=key[2],bit=bits[key],before=position,after=following))
                else:
                    trace_omitted+=1
                position=following
        return position

    for number,q in enumerate(queries):
        before=work
        position=q['index']
        hops=0
        try:
            while True:
                hops+=1
                position=apply_once(position,q['direction'],number,hops)
                if position<m:
                    break
                if hops>n-m:
                    raise ArithmeticError('Cycle restriction exceeded permutation path bound')
        except KeyError as exc:
            outputs.append(dict(query=q,status='missing_shared_bit',output=None,missing_key=list(exc.args[0]),switch_lookups=work-before))
        except OverflowError as exc:
            outputs.append(dict(query=q,status='budget_exceeded',output=None,reason=str(exc),switch_lookups=work-before))
        else:
            outputs.append(dict(query=q,status='complete',output=position,cycle_hops=hops,switch_lookups=work-before))
    return dict(status='complete' if all(r['status']=='complete' for r in outputs) else 'partial',
                mode='replay',dimension=d,full_domain_size=n,target_size=m,sweeps=sweeps,
                outputs=outputs,work_units=work,distinct_shared_bits_accessed=len(accessed),
                trace=trace,trace_records_omitted=trace_omitted,
                one_full_domain_query_lookup_count=sweeps*d,
                cycle_restriction_worst_case_full_permutation_calls=n-m+1,
                input_randomness='Explicit deterministic bit values; independent fair-bit distribution not established',
                storage_model='Sparse supplied bit map and bounded query traces; no full permutation table constructed',
                interpretation='Forward/inverse replay with one immutable bit per shared switch. Cycle restriction is a conventional additional construction and can require many full-domain calls. Missing bits and budget failures remain unknown; no approximate-uniform or cryptographic guarantee is granted.')


def exact_law(d,sweeps,work_budget):
    """Layerwise integer convolution over every fair layer assignment."""
    n=1<<d
    states={tuple(range(n)):1}
    work=0
    denominator=1
    for _ in range(sweeps):
        for coordinate in range(d):
            following=defaultdict(int)
            bit_position=d-1-coordinate
            for permutation,count in states.items():
                for choices in itertools.product((0,1),repeat=n//2):
                    image=[]
                    for position in permutation:
                        if work>=work_budget:
                            return None,denominator,work,'Permutation-image evaluation budget exhausted'
                        work+=1
                        image.append(position^((1<<bit_position) if choices[pair_id(position,bit_position)] else 0))
                    following[tuple(image)]+=count
                    if len(following)>MAX_STATES:
                        return None,denominator,work,'Exact permutation state cap exceeded'
            states=dict(following)
            denominator*=1<<(n//2)
    return states,denominator,work,None


def law(record,work_budget):
    if set(record)!={'mode','dimension','sweeps','target_size'}:
        raise ValueError('Invalid law fields')
    d=integer(record['dimension'],1,3,'exact law dimension')
    sweeps=integer(record['sweeps'],0,8,'sweeps')
    n=1<<d
    m=integer(record['target_size'],1,n,'target_size')
    states,denominator,work,reason=exact_law(d,sweeps,work_budget)
    if states is None:
        return dict(status='budget_exceeded',mode='exact_law',dimension=d,sweeps=sweeps,target_size=m,work_units=work,total_variation=None,permutation_law=None,reason=reason,
                    interpretation='The complete law is unknown; partial convolution is not returned as a probability certificate.')
    projected=defaultdict(int)
    for permutation,count in states.items():
        projected[induced(permutation,m)]+=count
    assert sum(projected.values())==denominator
    uniform=Fraction(1,factorial(m))
    tv=(sum((abs(Fraction(count,denominator)-uniform) for count in projected.values()),Fraction(0))+(factorial(m)-len(projected))*uniform)/2
    marginals=[[Fraction(0) for _ in range(m)] for _ in range(m)]
    for permutation,count in projected.items():
        for i,value in enumerate(permutation):
            marginals[i][value]+=Fraction(count,denominator)
    return dict(status='complete',mode='exact_law',dimension=d,sweeps=sweeps,full_domain_size=n,target_size=m,
                work_units=work,full_permutation_support=len(states),target_permutation_support=len(projected),
                fair_switch_bits=sweeps*d*n//2,coin_string_count=denominator,total_variation=str(tv),
                all_coordinate_marginals_uniform=all(v==Fraction(1,m) for row in marginals for v in row),
                coordinate_marginals=[[str(v) for v in row] for row in marginals],
                permutation_law=[dict(permutation=list(p),mass=str(Fraction(count,denominator)),coin_strings=count) for p,count in sorted(projected.items())],
                input_randomness='Every full switch assignment weighted by independent fair bits in this finite mathematical reference',
                interpretation='Exact tiny coordinate-sweep law and optional induced-cycle pushforward, starting from identity. This does not choose a universal source sweep constant or establish arbitrary-size convergence or cryptographic security.')


def calculate(record,work_budget=MAX_WORK):
    if not isinstance(record,dict) or record.get('mode') not in ('replay','exact_law','seed_support_bound'):
        raise ValueError('Require replay, exact_law or seed_support_bound mode')
    integer(work_budget,1,MAX_WORK,'work budget')
    if record['mode']=='replay':
        result=replay(record,work_budget)
    elif record['mode']=='exact_law':
        result=law(record,work_budget)
    else:
        if set(record)!={'mode','dimension','seed_bits'}:
            raise ValueError('Invalid seeded-support fields')
        d=integer(record['dimension'],1,63,'dimension')
        b=integer(record['seed_bits'],0,4096,'seed bits')
        n=1<<d
        # At most 2^b output permutations when all output randomness comes
        # from one b-bit seed. The remaining fixed inputs add no entropy.
        lower_bits=(n//2)*(d-1)
        result=dict(status='complete',mode='seed_support_bound',dimension=d,domain_size=n,seed_bits=b,
                    model='Deterministic full permutation of fixed domain, depending only on one b-bit seed; no extra independent randomness',
                    half_factorial_product_log2_lower_bound=lower_bits,
                    support_ratio_upper_bound_exponent=b-lower_bits,
                    conservative_distance_lower_bound_expression=f'1 - 2^({b-lower_bits})' if b<lower_bits else '0',
                    interpretation='A support lower bound on statistical total variation from uniform; not a computational-security claim, quality test or objection to seeded application goals that do not require this statistical guarantee.')
        if n<=128:
            size=factorial(n)
            result.update(exact_factorial=str(size),factorial_bit_length=size.bit_length(),
                          exact_distance_lower_bound=str(max(Fraction(0),1-Fraction(1<<b,size))))
    return dict(result,input_record_sha256=hashlib.sha256(json.dumps(record,sort_keys=True,separators=(',',':')).encode()).hexdigest(),
                source_connection=dict(family='238',source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',source_proof_verification='not_run'),
                formal_program_verification='not_run',limits=dict(work_budget=work_budget,exact_law_states=MAX_STATES,replay_trace_records=MAX_TRACE,wall_time_deadline='none'))


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