#!/usr/bin/env python3
"""Bounded exact positive-integer subset witnesses/counts, not source backends."""
import argparse
import bisect
import hashlib
import json
import re
from pathlib import Path

MAX_ITEMS = 64
MAX_RECORDS = 50_000
MAX_WORK = 2_000_000
MODEL = dict(value_model='positive_exact_integers', selection_model='each_distinct_item_at_most_once',
             target_model='one_exact_nonnegative_integer', additional_constraints='none', unit_model='one_declared_integer_unit')
METHODS = ('distinct_sum_MITM', 'occurrence_MITM')


class BudgetExceeded(Exception):
    pass


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


def parse(record):
    required = {'items', 'target', 'model'}
    if not isinstance(record, dict) or not required <= set(record) or set(record) - required - {'proposed_subset'}:
        raise ValueError('Invalid subset-sum schema')
    if not isinstance(record['model'], dict) or record['model'] != MODEL:
        raise ValueError('Require the exact positive-integer selection model')
    items = record['items']
    if not isinstance(items, list) or not 2 <= len(items) <= MAX_ITEMS:
        raise ValueError('Require two through sixty-four explicitly identified items')
    ids, values = [], []
    for item in items:
        if not isinstance(item, dict) or set(item) != {'id', 'value'}:
            raise ValueError('Invalid item record')
        identity = item['id']
        if not isinstance(identity, str) or not re.fullmatch(r'[A-Za-z0-9_-]{1,64}', identity) or identity in ids:
            raise ValueError('Invalid or repeated item identity')
        ids.append(identity)
        values.append(integer(item['value'], 1, 2**256 - 1, 'positive item value'))
    target = integer(record['target'], 0, 2**256 - 1, 'target')
    return ids, values, target


def audit_subset(ids, values, target, subset):
    if not isinstance(subset, list) or len(subset) > len(ids):
        return dict(valid=False, exact_sum=None, errors=['Invalid subset shape or size'])
    known = dict(zip(ids, values))
    seen, errors, total = set(), [], 0
    for identity in subset:
        if not isinstance(identity, str) or identity not in known:
            errors.append('Unknown item identity')
        elif identity in seen:
            errors.append('Repeated item identity ' + identity)
        else:
            seen.add(identity)
            total += known[identity]
    if not errors and total != target:
        errors.append('Exact integer sum differs from target')
    return dict(valid=not errors, exact_sum=total if not errors else None, errors=errors)


def calculate(record, method='distinct_sum_MITM', record_budget=MAX_RECORDS, work_budget=MAX_WORK):
    if method not in METHODS:
        raise ValueError('Unknown reference method')
    integer(record_budget, 1, MAX_RECORDS, 'stored record budget')
    integer(work_budget, 1, MAX_WORK, 'work budget')
    ids, values, target = parse(record)
    n, work, peak = len(ids), 0, 0
    proposed = audit_subset(ids, values, target, record['proposed_subset']) if 'proposed_subset' in record else None
    base = dict(status='unknown', exact_feasible=None, exact_index_subset_count=None, witness=None,
                method=method, n=n, target=target, proposed_subset_audit=proposed,
                limits=dict(items=MAX_ITEMS, stored_sum_records=record_budget, generation_and_probe_work=work_budget),
                source_connection=dict(family='138', source_commit='fd4aeeb2ee4fc729c18d98444fed42fd0529eeeb',
                                       faster_source_algorithm_implemented=False, lower_space_source_algorithm_implemented=False),
                source_proof_verification='not_run', formal_program_verification='not_run',
                input_sha256=hashlib.sha256(json.dumps(record, sort_keys=True, separators=(',', ':')).encode()).hexdigest(),
                interpretation='Conventional exact disjoint-half enumeration with optional equal-sum compression. Counts preserve distinct original index subsets, including repeated values. Caps count sum records and generation/probe work, excluding input/runtime objects and native sorting scratch/comparisons; they are not byte, wall-time or word-RAM guarantees. No signed amounts, tolerance, dates, mandatory choices, priorities or ledger actions are modeled.')
    if target == 0:
        return dict(base, status='complete', exact_feasible=True, exact_index_subset_count=1, witness=[], work=0, peak_stored_sum_records=0,
                    completion_basis='positive_inputs_have_unique_empty_zero_sum')
    if target > sum(values):
        return dict(base, status='complete', exact_feasible=False, exact_index_subset_count=0, witness=None, work=0, peak_stored_sum_records=0,
                    completion_basis='target_exceeds_total_positive_value')

    def tick():
        nonlocal work
        work += 1
        if work > work_budget:
            raise BudgetExceeded('Generation/probe work budget exhausted')

    def capacity(count):
        nonlocal peak
        if count > record_budget:
            raise BudgetExceeded('Stored sum-record budget exhausted')
        peak = max(peak, count)

    def generate(indices, retained):
        if method == 'distinct_sum_MITM':
            old = {0: (0, 1)}
            capacity(retained + 1)
            for i in indices:
                new = {}
                for total, (mask, multiplicity) in old.items():
                    for weight, witness_mask in ((total, mask), (total + values[i], mask | (1 << i))):
                        tick()
                        if weight > target:
                            continue
                        if weight in new:
                            previous_mask, previous_count = new[weight]
                            new[weight] = (previous_mask, previous_count + multiplicity)
                        else:
                            capacity(retained + len(old) + len(new) + 1)
                            new[weight] = (witness_mask, multiplicity)
                old = new
            return old
        old = [(0, 0)]
        capacity(retained + 1)
        for i in indices:
            new = []
            for total, mask in old:
                for weight, witness_mask in ((total, mask), (total + values[i], mask | (1 << i))):
                    tick()
                    if weight > target:
                        continue
                    capacity(retained + len(old) + len(new) + 1)
                    new.append((weight, witness_mask))
            old = new
        return old

    try:
        split = n // 2
        left = generate(range(split), 0)
        right = generate(range(split, n), len(left))
        count, witness_mask = 0, None
        if method == 'distinct_sum_MITM':
            for weight, (mask, multiplicity) in right.items():
                tick()
                match = left.get(target - weight)
                if match is not None:
                    left_mask, left_count = match
                    count += multiplicity * left_count
                    if witness_mask is None:
                        witness_mask = left_mask | mask
        else:
            left.sort()
            for weight, mask in right:
                tick()
                lo = bisect.bisect_left(left, (target - weight, -1))
                hi = bisect.bisect_left(left, (target - weight + 1, -1))
                count += hi - lo
                if witness_mask is None and lo < hi:
                    witness_mask = left[lo][1] | mask
    except BudgetExceeded as exc:
        return dict(base, status='budget_exceeded', reason=str(exc), work=work, peak_stored_sum_records=peak)
    witness = [identity for i, identity in enumerate(ids) if witness_mask is not None and witness_mask & (1 << i)] if count else None
    if count and not audit_subset(ids, values, target, witness)['valid']:
        raise ArithmeticError('Internal exact witness invalid')
    return dict(base, status='complete', exact_feasible=bool(count), exact_index_subset_count=count, witness=witness,
                work=work, peak_stored_sum_records=peak, left_final_records=len(left), right_final_records=len(right),
                completion_basis='complete_disjoint_half_enumeration_with_multiplicity')


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', type=Path)
    parser.add_argument('--method', choices=METHODS, default=METHODS[0])
    parser.add_argument('--record-budget', type=int, default=MAX_RECORDS)
    parser.add_argument('--work-budget', type=int, default=MAX_WORK)
    args = parser.parse_args()
    result = calculate(json.loads(args.input.read_text()), args.method, args.record_budget, args.work_budget)
    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()
