#!/usr/bin/env python3
"""Bounded exact perfect-matching counts, rank replay and edge-failure sensitivity.

Conventional subset dynamic programming; not the source FPRAS or entropy
optimizer. Counts labeled edge sets in a simple unweighted graph. Exhausted
work returns unknown, never a fabricated zero. No weighted/fairness claim.
"""
import argparse
import hashlib
import json
from fractions import Fraction
from pathlib import Path

MAX_VERTICES = 24
MAX_STATES = 200_000
MAX_WORK = 1_000_000
MAX_RANKS = 100


class BudgetExceeded(Exception):
    pass


def calculate(record, ranks=None, state_budget=MAX_STATES, work_budget=MAX_WORK):
    if not isinstance(record, dict) or set(record) != {'vertices', 'edges'}:
        raise ValueError('Provide an object with vertices and edges only')
    n = record['vertices']
    if isinstance(n, bool) or not isinstance(n, int) or not 0 <= n <= MAX_VERTICES:
        raise ValueError(f'vertices must be from 0 to {MAX_VERTICES}')
    for label, value, maximum in [('state_budget', state_budget, MAX_STATES), ('work_budget', work_budget, MAX_WORK)]:
        if isinstance(value, bool) or not isinstance(value, int) or not 1 <= value <= maximum:
            raise ValueError(f'{label} must be from 1 to {maximum}')
    raw_edges = record['edges']
    if not isinstance(raw_edges, list) or len(raw_edges) > n * (n - 1) // 2:
        raise ValueError('edges must be a bounded list of undirected pairs')
    edges = set()
    adjacency = [0] * n
    for edge in raw_edges:
        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('Each edge must be a pair of integer vertex indices')
        a, b = sorted(edge)
        if not 0 <= a < b < n or (a, b) in edges:
            raise ValueError('Edges must be distinct, loopless and in range')
        edges.add((a, b))
        adjacency[a] |= 1 << b
        adjacency[b] |= 1 << a
    ranks = [] if ranks is None else ranks
    if (not isinstance(ranks, list) or len(ranks) > MAX_RANKS
            or any(isinstance(r, bool) or not isinstance(r, int) or r < 0 for r in ranks)):
        raise ValueError(f'ranks must be at most {MAX_RANKS} nonnegative integers')
    memo = {}
    work = 0

    def charge():
        nonlocal work
        if work >= work_budget:
            raise BudgetExceeded('subset/neighbor work budget exhausted')
        work += 1

    def count(mask):
        if mask in memo:
            return memo[mask]
        charge()
        if len(memo) >= state_budget:
            raise BudgetExceeded('memo-state budget exhausted')
        if mask == 0:
            answer = 1
        elif mask.bit_count() % 2:
            answer = 0
        else:
            first = mask & -mask
            a = first.bit_length() - 1
            rest = mask ^ first
            choices = adjacency[a] & rest
            answer = 0
            while choices:
                charge()
                partner = choices & -choices
                choices ^= partner
                answer += count(rest ^ partner)
        # Recursion may have filled the memo after this call was started.
        if len(memo) >= state_budget:
            raise BudgetExceeded('memo-state budget exhausted')
        memo[mask] = answer
        return answer

    all_vertices = (1 << n) - 1
    result = dict(
        input_record_sha256=hashlib.sha256(json.dumps(record, sort_keys=True, separators=(',', ':')).encode()).hexdigest(),
        vertices=n, edges=len(edges), objective='number_of_labeled_unweighted_perfect_matchings',
        method='conventional_exact_subset_dynamic_program',
        source_connection=dict(family='113', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                               source_fpras_implemented=False, source_proof_verification='not_run'),
        formal_program_verification='not_run',
        limits=dict(vertices=MAX_VERTICES, memo_states=state_budget, count_work_units=work_budget,
                    requested_ranks=MAX_RANKS, wall_time_deadline='none'),
        work_unit_definition='New subset evaluations plus partner-branch visits; not bit operations or measured time.')
    try:
        total = count(all_vertices)
    except BudgetExceeded as exc:
        return dict(result, status='budget_exceeded', exact_count=None, feasible=None,
                    reason=str(exc), memo_states=len(memo), work_units=work,
                    edge_sensitivity=None, ranked_matchings=[])
    if any(rank >= total for rank in ranks):
        raise ValueError('Each requested rank must be smaller than the exact count')
    sensitivity = []
    try:
        for a, b in sorted(edges):
            # Every perfect matching containing ab is uniquely its remainder.
            contains = count(all_vertices ^ (1 << a) ^ (1 << b)) if n % 2 == 0 else 0
            sensitivity.append(dict(edge=[a, b], containing_matchings=contains,
                                    surviving_matchings_if_edge_removed=total-contains,
                                    uniform_edge_marginal=str(Fraction(contains, total)) if total else None,
                                    uniform_survival_fraction=str(Fraction(total-contains, total)) if total else None,
                                    forced_when_feasible=bool(total and contains == total),
                                    belongs_to_any_perfect_matching=bool(contains)))
    except BudgetExceeded as exc:
        return dict(result, status='count_complete_sensitivity_budget_exceeded', exact_count=total,
                    feasible=bool(total), reason=str(exc), memo_states=len(memo), work_units=work,
                    edge_sensitivity=sensitivity, ranked_matchings=[])
    ranked = []
    for original_rank in ranks:
        rank = original_rank
        mask = all_vertices
        matching = []
        # All these suffix counts were evaluated during the complete count.
        while mask:
            first = mask & -mask
            a = first.bit_length() - 1
            rest = mask ^ first
            choices = adjacency[a] & rest
            while choices:
                partner = choices & -choices
                choices ^= partner
                branch = memo[rest ^ partner]
                if rank < branch:
                    matching.append([a, partner.bit_length()-1])
                    mask = rest ^ partner
                    break
                rank -= branch
            else:
                raise AssertionError('Rank did not select a feasible branch')
        ranked.append(dict(rank=original_rank, matching=matching))
    return dict(result, status='complete', exact_count=total, feasible=bool(total),
                memo_states=len(memo), work_units=work, edge_sensitivity=sensitivity,
                ranked_matchings=ranked,
                interpretation='Exact for this bounded graph. Edge-failure counts measure surviving unweighted perfect pairings, not weighted plan quality, fairness or resilience to arbitrary failures. Uniform rank would induce a uniform matching law; this utility does not draw random ranks.')


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


if __name__ == '__main__':
    main()
