#!/usr/bin/env python3
"""Bounded exact counting and rank-based generation of fixed-margin tables.

This conventional row dynamic program is a small reference, not the new
polynomial-bit sampler/FPRAS. Uniform ranks correspond bijectively to uniform
tables; OS-entropy draws use secrets.randbelow. This is not the conditional
independence distribution and provides no privacy guarantee.
"""
import argparse
import json
import secrets
from pathlib import Path

MAX_DIMENSION = 6
MAX_TOTAL = 128
MAX_STATES = 50_000
MAX_COUNT_WORK = 1_000_000
MAX_SAMPLE_WORK = 500_000


class BudgetExceeded(RuntimeError):
    pass


def integer_vector(value, name):
    if not isinstance(value, list) or not 1 <= len(value) <= MAX_DIMENSION:
        raise ValueError(f'{name} must have 1 to {MAX_DIMENSION} entries')
    if any(isinstance(x, bool) or not isinstance(x, int) or x < 0 for x in value):
        raise ValueError(f'{name} must contain nonnegative integers')
    return tuple(value)


class TableSpace:
    def __init__(self, record, state_budget=MAX_STATES, count_work_budget=MAX_COUNT_WORK):
        for value, maximum, name in [(state_budget,MAX_STATES,'state budget'), (count_work_budget,MAX_COUNT_WORK,'work budget')]:
            if isinstance(value, bool) or not isinstance(value, int) or not 1 <= value <= maximum:
                raise ValueError(f'{name} must be an integer from 1 to {maximum}')
        if not isinstance(record, dict):
            raise ValueError('Input must be an object')
        self.rows = integer_vector(record.get('rows'), 'rows')
        self.columns = integer_vector(record.get('columns'), 'columns')
        self.total = sum(self.rows)
        if self.total != sum(self.columns):
            raise ValueError('Margins must have equal totals')
        if self.total > MAX_TOTAL:
            raise ValueError(f'This reference is limited to total {MAX_TOTAL}')
        self.m, self.n = len(self.rows), len(self.columns)
        supplied = record.get('cell_bounds')
        if supplied is None:
            self.bounds = tuple(tuple(min(r, c) for c in self.columns) for r in self.rows)
        else:
            if not isinstance(supplied, list) or len(supplied) != self.m:
                raise ValueError('cell_bounds must have one row per margin row')
            checked = [integer_vector(row, 'cell bound row') for row in supplied]
            if any(len(row) != self.n for row in checked):
                raise ValueError('cell_bounds shape must match margins')
            self.bounds = tuple(tuple(min(value, self.rows[i], self.columns[j]) for j, value in enumerate(row))
                                for i, row in enumerate(checked))
        self.future_capacity = [None]*(self.m+1)
        self.future_capacity[self.m] = (0,)*self.n
        for i in reversed(range(self.m)):
            self.future_capacity[i] = tuple(self.bounds[i][j]+self.future_capacity[i+1][j] for j in range(self.n))
        self.memo = {}
        self.states = 0
        self.work = 0
        self.work_limit = count_work_budget
        self.state_limit = state_budget

    def tick(self):
        self.work += 1
        if self.work > self.work_limit:
            raise BudgetExceeded('Reference work budget exceeded')

    def row_options(self, row_index, residual):
        caps = [min(self.bounds[row_index][j], residual[j]) for j in range(self.n)]
        suffix = [0]*(self.n+1)
        for j in reversed(range(self.n)):
            suffix[j] = suffix[j+1]+caps[j]
        chosen = [0]*self.n
        def visit(j, remaining):
            self.tick()
            if j == self.n:
                if remaining == 0:
                    yield tuple(chosen)
                return
            lower = max(0, remaining-suffix[j+1])
            upper = min(caps[j], remaining)
            for value in range(lower, upper+1):
                chosen[j] = value
                yield from visit(j+1, remaining-value)
        yield from visit(0, self.rows[row_index])

    def count(self, row_index=0, residual=None):
        if residual is None:
            residual = self.columns
        key = (row_index, residual)
        if key in self.memo:
            return self.memo[key]
        self.tick()
        self.states += 1
        if self.states > self.state_limit:
            raise BudgetExceeded('Reference memo-state budget exceeded')
        if row_index == self.m:
            answer = int(all(value == 0 for value in residual))
        elif any(residual[j] > self.future_capacity[row_index][j] for j in range(self.n)):
            answer = 0
        else:
            answer = 0
            for row in self.row_options(row_index, residual):
                following = tuple(residual[j]-row[j] for j in range(self.n))
                answer += self.count(row_index+1, following)
        self.memo[key] = answer
        return answer

    def unrank(self, rank):
        total = self.count()
        if isinstance(rank, bool) or not isinstance(rank, int) or not 0 <= rank < total:
            raise ValueError('Rank must be an integer in [0, exact_count)')
        residual = self.columns
        table = []
        for i in range(self.m):
            for row in self.row_options(i, residual):
                following = tuple(residual[j]-row[j] for j in range(self.n))
                completions = self.count(i+1, following)
                if rank >= completions:
                    rank -= completions
                else:
                    table.append(list(row))
                    residual = following
                    break
            else:
                raise AssertionError('No row found for a validated rank')
        if rank != 0 or any(residual):
            raise AssertionError('Unranking invariant failed')
        return table


def calculate(record, ranks=None, random_samples=0, state_budget=MAX_STATES, count_work_budget=MAX_COUNT_WORK):
    ranks = [] if ranks is None else ranks
    if not isinstance(ranks, list) or len(ranks) > 100 or any(isinstance(x, bool) or not isinstance(x, int) or x < 0 for x in ranks):
        raise ValueError('Use up to 100 nonnegative integer ranks')
    if isinstance(random_samples, bool) or not isinstance(random_samples, int) or random_samples < 0 or len(ranks)+random_samples > 100:
        raise ValueError('Use at most 100 requested tables')
    space = TableSpace(record, state_budget, count_work_budget)
    result = dict(status='counting', rows=list(space.rows), columns=list(space.columns),
                  effective_cell_bounds=[list(row) for row in space.bounds], total=space.total,
                  exact_count=None, feasible=None, tables=[],
                  arithmetic='exact_python_integers', algorithm='bounded_row_dynamic_program_and_lexicographic_unranking',
                  probability_law='uniform_over_feasible_tables_if_rank_is_uniform',
                  random_source='secrets_randbelow_OS_entropy' if random_samples else 'caller_supplied_ranks',
                  limits=dict(maximum_dimension=MAX_DIMENSION, maximum_total=MAX_TOTAL,
                              memo_states=state_budget, count_work=count_work_budget, sample_work=MAX_SAMPLE_WORK),
                  research_connection=dict(family='115', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                                           new_polynomial_bit_backend_implemented=False),
                  independent_proof_verification='not_run', privacy_guarantee=False,
                  interpretation='This finite reference counts tables once each. Uniform tables differ from the conditional independence distribution. Budget exhaustion is unknown, never zero or infeasibility.')
    try:
        total = space.count()
    except BudgetExceeded as error:
        result.update(status='budget_exceeded', reason=str(error), memo_states=space.states, count_work=space.work)
        return result
    result.update(exact_count=total, feasible=total > 0, memo_states=space.states, count_work=space.work)
    if any(rank >= total for rank in ranks):
        raise ValueError('Requested rank lies outside the feasible table space')
    if total == 0:
        result.update(status='empty', sample_work=0)
        return result
    requested = [(rank, 'supplied') for rank in ranks]
    requested += [(secrets.randbelow(total), 'OS_entropy_draw') for _ in range(random_samples)]
    space.work, space.work_limit = 0, MAX_SAMPLE_WORK
    try:
        for rank, selection in requested:
            result['tables'].append(dict(rank=rank, selection=selection, table=space.unrank(rank)))
    except BudgetExceeded as error:
        result.update(status='count_complete_sampling_budget_exceeded', reason=str(error), sample_work=space.work)
        return result
    result.update(status='complete', sample_work=space.work)
    return result


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('--random-samples', type=int, default=0)
    parser.add_argument('--output', type=Path)
    args = parser.parse_args()
    result = json.dumps(calculate(json.loads(args.input.read_text()), args.rank, args.random_samples), 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()
