#!/usr/bin/env python3
"""Enumerate a bounded degree fiber and audit the precise lazy switch kernel.

Exact integer propagation with a shared rational denominator. Optional host
constraints are experimental and receive no complete-host source guarantee.
This does not implement the source's unrestricted exact residual sampler.
"""
import argparse
import itertools
import json
from collections import Counter
from fractions import Fraction
from math import comb
from pathlib import Path

MAX_VERTICES = 6
MAX_STATES = 512
MAX_STEPS = 64
MAX_WORK = 2_000_000


class BudgetExceeded(Exception):
    pass


def audit(record, steps=12, work_budget=MAX_WORK):
    if not isinstance(record, dict) or set(record)-{'degrees','allowed_edges'} or 'degrees' not in record:
        raise ValueError('Provide degrees and optionally allowed_edges only')
    degrees=record['degrees']
    if not isinstance(degrees,list) or len(degrees)>MAX_VERTICES:
        raise ValueError(f'degrees must be a list with at most {MAX_VERTICES} entries')
    n=len(degrees)
    if any(isinstance(d,bool) or not isinstance(d,int) or not 0<=d<n for d in degrees):
        raise ValueError('Degree entries must be integer values from 0 to n-1')
    if isinstance(steps,bool) or not isinstance(steps,int) or not 0<=steps<=MAX_STEPS:
        raise ValueError(f'steps must be from 0 to {MAX_STEPS}')
    if isinstance(work_budget,bool) or not isinstance(work_budget,int) or not 1<=work_budget<=MAX_WORK:
        raise ValueError(f'work_budget must be from 1 to {MAX_WORK}')
    pairs=list(itertools.combinations(range(n),2))
    pair_index={pair:i for i,pair in enumerate(pairs)}
    host=(1<<len(pairs))-1
    if 'allowed_edges' in record:
        allowed=record['allowed_edges']
        if not isinstance(allowed,list) or len(allowed)>len(pairs):
            raise ValueError('allowed_edges must be a bounded list of pairs')
        host=0
        for edge in allowed:
            if not isinstance(edge,list) or len(edge)!=2 or any(isinstance(v,bool) or not isinstance(v,int) for v in edge):
                raise ValueError('Allowed edges must be integer vertex pairs')
            pair=tuple(sorted(edge))
            if pair not in pair_index:
                raise ValueError('Allowed edge must be loopless and in range')
            bit=1<<pair_index[pair]
            if host&bit:
                raise ValueError('Duplicate allowed edge')
            host|=bit
    complete_host=host==(1<<len(pairs))-1
    work=0
    def charge():
        nonlocal work
        if work>=work_budget:
            raise BudgetExceeded('enumeration/proposal/propagation budget exhausted')
        work+=1
    result=dict(degrees=degrees,vertices=n,steps=steps,complete_host=complete_host,
                source_connection=dict(family='131',source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                    selected_chain_model_applicable=complete_host and n>=4,
                    source_proof_verification='not_run',unrestricted_exact_source_sampler_implemented=False),
                limits=dict(vertices=MAX_VERTICES,states=MAX_STATES,steps=MAX_STEPS,work_units=work_budget),
                arithmetic='exact integers and fractions',formal_program_verification='not_run',
                work_unit_definition='Candidate graph tests, selected-edge degree updates, kernel proposals and nonzero distribution transitions; not bit operations or wall time.')
    try:
        states=[]
        host_bits=[1<<i for i in range(len(pairs)) if host&(1<<i)]
        edge_target=sum(degrees)//2 if sum(degrees)%2==0 else -1
        for subset in range(1<<len(host_bits)):
            charge()
            if subset.bit_count()!=edge_target:
                continue
            mask=0
            for i,bit in enumerate(host_bits):
                if subset&(1<<i):
                    mask|=bit
            actual=[0]*n
            remaining=mask
            while remaining:
                charge()
                bit=remaining&-remaining
                remaining^=bit
                a,b=pairs[bit.bit_length()-1]
                actual[a]+=1;actual[b]+=1
            if actual==degrees:
                if len(states)>=MAX_STATES:
                    raise BudgetExceeded('state cap exhausted')
                states.append(mask)
        if not states:
            return dict(result,status='complete_infeasible',state_count=0,feasible=False,work_units=work,
                        interpretation='Exhaustive bounded host enumeration found no labeled realization.')
        states.sort()
        indices={mask:i for i,mask in enumerate(states)}
        q=12*comb(n,4) if n>=4 else 1
        proposals=[]
        for a,b,c,d in itertools.combinations(range(n),4):
            matchings=[[(a,b),(c,d)],[(a,c),(b,d)],[(a,d),(b,c)]]
            masks=[sum(1<<pair_index[pair] for pair in matching) for matching in matchings]
            proposals.extend((f,g) for f in masks for g in masks if f!=g)
        rows=[]
        for i,mask in enumerate(states):
            transitions=Counter()
            for removed,added in proposals:
                charge()
                if mask&removed==removed and mask&added==0 and added&host==added:
                    changed=(mask^removed)|added
                    transitions[indices[changed]]+=1
            transitions[i]=q-sum(transitions.values())
            rows.append(dict(transitions))
        symmetric=all(rows[j].get(i,0)==weight for i,row in enumerate(rows) for j,weight in row.items())
        stochastic=all(sum(row.values())==q and row.get(i,0)*2>=q for i,row in enumerate(rows))
        if not symmetric or not stochastic:
            raise AssertionError('Kernel convention or finite state-space consistency failed')
        seen=set();components=[]
        for start in range(len(states)):
            if start in seen:continue
            stack=[start];seen.add(start);component=[]
            while stack:
                current=stack.pop();component.append(current)
                for neighbor,weight in rows[current].items():
                    if neighbor!=current and weight and neighbor not in seen:
                        seen.add(neighbor);stack.append(neighbor)
            components.append(sorted(component))
        m=len(states)
        numerators=[0]*m;numerators[0]=1;denominator=1
        tv=[]
        first_quarter=None
        for time in range(steps+1):
            distance=Fraction(sum(abs(value*m-denominator) for value in numerators),2*m*denominator)
            tv.append(dict(step=time,total_variation=str(distance)))
            if distance<=Fraction(1,4) and first_quarter is None:first_quarter=time
            if time==steps:break
            updated=[0]*m
            for i,value in enumerate(numerators):
                if not value:continue
                for j,weight in rows[i].items():
                    charge()
                    updated[j]+=value*weight
            numerators=updated;denominator*=q
        neighbor_counts=[len(row)-1 for row in rows]
        sum_neighbors=sum(neighbor_counts)
        jump_stationary=None
        if len(components)==1 and m>1:
            jump_stationary=Fraction(sum(abs(Fraction(count,sum_neighbors)-Fraction(1,m)) for count in neighbor_counts),2)
        return dict(result,status='complete',feasible=True,state_count=m,work_units=work,
            transition_denominator=q,rows_symmetric=symmetric,rows_stochastic_and_lazy=stochastic,
            switch_component_sizes=[len(c) for c in components],switch_connected=len(components)==1,
            starting_state_index=0,starting_state_edges=[list(pair) for i,pair in enumerate(pairs) if states[0]&(1<<i)],
            total_variation_by_step=tv,first_quarter_tv_from_selected_start=first_quarter,
            source_quarter_tv_bound=2*n**8 if complete_host and n>=4 else None,
            valid_neighbor_count_histogram=dict(sorted(Counter(neighbor_counts).items())),
            generic_accepted_switch_jump_chain=dict(
                stationary_law='proportional to valid-switch neighbor count in a connected nontrivial fiber',
                exact_stationary_tv_from_uniform=str(jump_stationary) if jump_stationary is not None else None,
                scope='Generic chain selecting uniformly among distinct valid switch neighbors. This is not a complete audit of any external library implementation.'),
            finite_states=[dict(index=i,edges=[list(pair) for j,pair in enumerate(pairs) if mask&(1<<j)]) for i,mask in enumerate(states)],
            interpretation='Finite exhaustive kernel and one-start law calculation. Extra host constraints can disconnect the chain. Empirical/fixed-start convergence does not verify the source theorem or worst-case large-graph mixing; exact unrestricted residual sampling is not implemented.')
    except BudgetExceeded as exc:
        return dict(result,status='budget_exceeded',feasible=None,state_count=None,work_units=work,
                    reason=str(exc),interpretation='No complete fiber/kernel conclusion is granted after budget exhaustion.')


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input',type=Path)
    parser.add_argument('--steps',type=int,default=12)
    parser.add_argument('--output',type=Path)
    args=parser.parse_args()
    body=json.dumps(audit(json.loads(args.input.read_text()),args.steps),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()
