#!/usr/bin/env python3
"""Exact bounded fixed-margin law/event comparison, never a new sampler.

Uniform integer tables and fixed-margin conditional independence are distinct
models. Bounds require an explicitly named conditioned independence model.
The intended event is fixed across laws, including probability-order events.
No privacy, causal, model-design or source-FPRAS certification is granted.
"""
import argparse
import hashlib
import json
from collections import defaultdict
from fractions import Fraction
from math import factorial
from pathlib import Path

from contingency_reference import TableSpace, BudgetExceeded, MAX_STATES, MAX_COUNT_WORK

MAX_TABLES = 2_000
MAX_ENUM_WORK = 4_000_000
LAWS = ('uniform_tables', 'conditional_independence', 'conditional_independence_given_bounds')


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


def matrix(value, m, n, lower, upper, name):
    if not isinstance(value, list) or len(value) != m or any(not isinstance(row, list) or len(row) != n for row in value):
        raise ValueError('Invalid ' + name + ' shape')
    return [[integer(x, lower, upper, name) for x in row] for row in value]


def rational(value):
    if not isinstance(value, str) or len(value) > 64:
        raise ValueError('Require canonical rational threshold string')
    try:
        result = Fraction(value)
    except (ValueError, ZeroDivisionError):
        raise ValueError('Invalid rational threshold') from None
    if str(result) != value or not 0 < result < 1:
        raise ValueError('Threshold must be canonical and strictly between zero and one')
    return result


def calculate(record, state_budget=MAX_STATES, count_work_budget=MAX_COUNT_WORK, enum_work_budget=MAX_ENUM_WORK, table_limit=MAX_TABLES):
    required = {'rows', 'columns', 'observed', 'intended_law', 'candidate_law', 'event', 'threshold'}
    if not isinstance(record, dict) or not required <= set(record) or set(record) - required - {'cell_bounds'}:
        raise ValueError('Invalid law-audit schema')
    integer(enum_work_budget, 1, MAX_ENUM_WORK, 'enumeration work budget')
    integer(table_limit, 1, MAX_TABLES, 'table cap')
    for name in ('intended_law', 'candidate_law'):
        law = record[name]
        if law not in LAWS:
            raise ValueError('Unknown law')
        if 'cell_bounds' in record and law == 'conditional_independence':
            raise ValueError('Bounded support requires conditional_independence_given_bounds')
        if 'cell_bounds' not in record and law == 'conditional_independence_given_bounds':
            raise ValueError('Conditioned law requires explicit cell bounds')
    alpha = rational(record['threshold'])
    space_record = {name: record[name] for name in ('rows', 'columns', 'cell_bounds') if name in record}
    space = TableSpace(space_record, state_budget, count_work_budget)
    observed = matrix(record['observed'], space.m, space.n, 0, space.total, 'observed table')
    if [sum(row) for row in observed] != list(space.rows) or [sum(observed[i][j] for i in range(space.m)) for j in range(space.n)] != list(space.columns):
        raise ValueError('Observed margins differ')
    if any(observed[i][j] > space.bounds[i][j] for i in range(space.m) for j in range(space.n)):
        raise ValueError('Observed table violates bounds')
    event = record['event']
    if not isinstance(event, dict):
        raise ValueError('Require explicit event')
    if event.get('kind') == 'linear_score_ge_observed' and set(event) == {'kind', 'coefficients'}:
        coefficients = matrix(event['coefficients'], space.m, space.n, -1_000_000, 1_000_000, 'score coefficients')
        observed_score = sum(coefficients[i][j] * observed[i][j] for i in range(space.m) for j in range(space.n))
    elif event == {'kind': 'probability_le_observed_under_intended_law'}:
        coefficients, observed_score = None, None
    else:
        raise ValueError('Unsupported event')
    base = dict(status='unknown', exact_table_count=None, distributions=None, event_comparison=None,
                total_variation=None, base_independence_constraint_event_probability=None,
                intended_law=record['intended_law'], candidate_law=record['candidate_law'],
                input_sha256=hashlib.sha256(json.dumps(record, sort_keys=True, separators=(',', ':')).encode()).hexdigest(),
                limits=dict(maximum_dimension=6, total=128, tables=table_limit, memo_states=state_budget,
                            count_work=count_work_budget, enumeration_work=enum_work_budget),
                source_connection=dict(family='115', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                                       unrestricted_polynomial_bit_sampler_or_FPRAS_implemented=False),
                source_proof_verification='not_run', formal_program_verification='not_run', privacy_guarantee=False,
                interpretation='Exact finite law comparison for a supplied design, not validation of that design. Bounds condition the ordinary fixed-margin independence law onto the specified feasible event; they do not automatically model real structural-zero mechanisms. Unknown budget status grants no partial probability or test decision.')
    try:
        count = space.count()
    except BudgetExceeded as exc:
        return dict(base, status='budget_exceeded', reason=str(exc), count_work=space.work, memo_states=space.states)
    base['exact_table_count'] = count
    count_work = space.work
    if count > table_limit:
        return dict(base, status='budget_exceeded', reason='Exact table count exceeds law enumeration cap', count_work=count_work, memo_states=space.states)
    if count == 0:
        # A validated observed table implies feasibility. Keep this invariant
        # distinct from an incomplete enumeration or a scientific inference.
        raise ArithmeticError('Validated observed table missing from feasible space')
    space.work, space.work_limit = 0, enum_work_budget
    tables, weights, scores = [], [], []
    facts = [factorial(i) for i in range(space.total + 1)]
    try:
        for rank in range(count):
            table = space.unrank(rank)
            divisor = 1
            for row in table:
                for cell in row:
                    space.tick()
                    divisor *= facts[cell]
            tables.append(table)
            weights.append(Fraction(1, divisor))
            if coefficients is not None:
                scores.append(sum(coefficients[i][j] * table[i][j] for i in range(space.m) for j in range(space.n)))
    except BudgetExceeded as exc:
        return dict(base, status='budget_exceeded', reason=str(exc), count_work=count_work, enumeration_work=space.work, memo_states=space.states)
    partition = sum(weights)
    independence = [weight / partition for weight in weights]
    uniform = [Fraction(1, count)] * count
    by_law = {'uniform_tables': uniform, 'conditional_independence': independence,
              'conditional_independence_given_bounds': independence}
    observed_rank = next(i for i, table in enumerate(tables) if table == observed)
    intended, candidate = by_law[record['intended_law']], by_law[record['candidate_law']]
    if coefficients is not None:
        selected = [i for i, score in enumerate(scores) if score >= observed_score]
        event_description = 'Declared linear score at least its observed value, held fixed across both laws'
    else:
        selected = [i for i, mass in enumerate(intended) if mass <= intended[observed_rank]]
        event_description = 'Tables at most as probable as observed under the intended law; the same selected set is evaluated under the candidate law'
    intended_tail, candidate_tail = sum(intended[i] for i in selected), sum(candidate[i] for i in selected)
    constant = Fraction(1, facts[space.total])
    for margin in space.rows + space.columns:
        constant *= facts[margin]
    constraint_mass = constant * partition
    score_law = None
    if coefficients is not None:
        histogram = defaultdict(lambda: [Fraction(0), Fraction(0)])
        for score, p, q in zip(scores, intended, candidate):
            histogram[score][0] += p; histogram[score][1] += q
        score_law = [dict(score=score, intended_mass=str(value[0]), candidate_mass=str(value[1])) for score, value in sorted(histogram.items())]
    return dict(base, status='complete', count_work=count_work, enumeration_work=space.work, memo_states=space.states,
                base_independence_constraint_event_probability=str(constraint_mass),
                total_variation=str(sum(abs(p - q) for p, q in zip(intended, candidate)) / 2),
                event_comparison=dict(definition=event_description, kind=event['kind'], observed_score=observed_score,
                                      observed_rank=observed_rank, event_ranks=selected, event_table_count=len(selected),
                                      intended_probability=str(intended_tail), candidate_probability=str(candidate_tail),
                                      absolute_error=str(abs(intended_tail - candidate_tail)), threshold=str(alpha),
                                      intended_at_or_below_threshold=intended_tail <= alpha,
                                      candidate_at_or_below_threshold=candidate_tail <= alpha,
                                      threshold_classification_changed=(intended_tail <= alpha) != (candidate_tail <= alpha)),
                score_distribution=score_law,
                distributions=[dict(rank=i, table=table, intended_mass=str(intended[i]), candidate_mass=str(candidate[i]),
                                    uniform_mass=str(uniform[i]), conditioned_independence_mass=str(independence[i])) for i, table in enumerate(tables)])


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', type=Path)
    parser.add_argument('--table-limit', type=int, default=MAX_TABLES)
    parser.add_argument('--enum-work-budget', type=int, default=MAX_ENUM_WORK)
    args = parser.parse_args()
    result = calculate(json.loads(args.input.read_text()), enum_work_budget=args.enum_work_budget, table_limit=args.table_limit)
    body = json.dumps(result, 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()
