#!/usr/bin/env python3
"""Declared three-community SBM scope, exact parameter envelopes and finite scores.

This audits a supplied model declaration and label scores. It does not fit or
authenticate a model, implement a threshold-optimal detector, or prove the
source theorem. Input/output coefficients use exact integer/rational strings.
"""
import argparse
import hashlib
import itertools
import json
import re
from fractions import Fraction
from pathlib import Path

MAX_BYTES = 1024**2
MAX_N = 10000
PERMUTATIONS = tuple(itertools.permutations(range(3)))


class Invalid(ValueError):
    pass


def rational(x):
    if type(x) is int:
        value = Fraction(x)
    elif type(x) is str and len(x) <= 80 and re.fullmatch(r'(0|[1-9][0-9]*)(/[1-9][0-9]*)?', x):
        value = Fraction(x)
    else:
        raise Invalid('use nonnegative integer or canonical rational strings')
    if value < 0 or max(value.numerator.bit_length(), value.denominator.bit_length()) > 128:
        raise Invalid('parameter must be nonnegative and at most 128 bits per component')
    return value


def integer(x, low, high, name):
    if type(x) is not int or not low <= x <= high:
        raise Invalid(f'{name}: integer from {low} through {high} required')
    return x


def interval(x):
    if type(x) is not list or len(x) != 2:
        raise Invalid('parameter interval must contain [lower, upper]')
    low, high = map(rational, x)
    if high < low:
        raise Invalid('interval endpoints are reversed')
    return low, high


def margin(a, b):
    return (a-b)**2-3*(a+2*b)


def margin_envelope(a, b):
    """Exact min/max on a rectangle: convex max at corners; min on edges.

    Both partial derivatives cannot vanish (their sum is -9), so include
    edge stationary points a=b+3/2 and b=a+3, together with all corners.
    """
    corners = [(x, y) for x in a for y in b]
    candidates = set(corners)
    for y in b:
        x = y+Fraction(3, 2)
        if a[0] <= x <= a[1]:
            candidates.add((x, y))
    for x in a:
        y = x+3
        if b[0] <= y <= b[1]:
            candidates.add((x, y))
    scored = [(margin(x, y), x, y) for x, y in sorted(candidates)]
    low = min(scored)
    high = max((margin(x, y), x, y) for x, y in corners)
    return dict(minimum=str(low[0]), maximum=str(high[0]),
                minimum_at=[str(low[1]), str(low[2])], maximum_at=[str(high[1]), str(high[2])],
                candidates=[dict(a=str(x), b=str(y), margin=str(v)) for v, x, y in scored])


def labels(x, name):
    if type(x) is not list or not 1 <= len(x) <= MAX_N:
        raise Invalid(f'{name}: expected 1 through {MAX_N} labels')
    for label in x:
        integer(label, 0, 2, name)
    return x


def score(truth, predicted):
    truth, predicted = labels(truth, 'truth'), labels(predicted, 'predicted')
    if len(truth) != len(predicted):
        raise Invalid('label lengths must match')
    counts = [[0]*3 for _ in range(3)]
    for i, j in zip(truth, predicted):
        counts[i][j] += 1
    agreements = [(sum(counts[i][perm[i]] for i in range(3)), perm) for perm in PERMUTATIONS]
    best, perm = max(agreements)
    n = len(truth)
    return dict(overlap=str(Fraction(best, n)), agreement_count=best, count=n,
                truth_to_prediction_permutation=list(perm), confusion=counts,
                unaligned_accuracy=str(Fraction(sum(counts[i][i] for i in range(3)), n)),
                constant_predictor_overlap=str(Fraction(max(map(sum, counts)), n)),
                interpretation='Finite permutation-aligned score; exceeding 1/3 alone is not evidence of asymptotic recovery')


def exact_uniform_label_baseline(truth):
    """Enumerate all independent uniform predictions for at most eight vertices."""
    truth = labels(truth, 'truth')
    if len(truth) > 8:
        return dict(status='unknown_size_limit', max_vertices=8)
    total = 0
    count = 3**len(truth)
    for predicted in itertools.product(range(3), repeat=len(truth)):
        total += max(sum(int(predicted[j] == perm[truth[j]]) for j in range(len(truth)))
                     for perm in PERMUTATIONS)
    return dict(status='exact', predictions_enumerated=count,
                expected_overlap=str(Fraction(total, count*len(truth))),
                randomization='All iid uniform predicted labelings, conditional on supplied truth')


def evaluate(doc):
    try:
        required = {'schema','communities','n','label_law','graph_kind','observation','a','b'}
        if type(doc) is not dict or not required <= set(doc) or set(doc)-required-{'evaluation'}:
            raise Invalid('unexpected or missing model fields')
        if doc['schema'] != 'community-recovery-audit-v1':
            raise Invalid('unknown schema')
        n = integer(doc['n'], 1, MAX_N, 'n')
        q = integer(doc['communities'], 2, 32, 'communities')
        a, b = interval(doc['a']), interval(doc['b'])
        for key in ['label_law','graph_kind','observation']:
            if type(doc[key]) is not str or len(doc[key]) > 64:
                raise Invalid(f'{key}: short declaration string required')
        reasons = []
        for condition, reason in [
                (q == 3, 'Corollary 1.2 is for three communities'),
                (doc['label_law'] == 'iid_uniform', 'Independent uniform labels required; fixed sizes are a different finite ensemble'),
                (doc['graph_kind'] == 'simple_undirected_independent_edges', 'Simple undirected conditionally independent-edge graph required'),
                (doc['observation'] == 'graph_only', 'Estimator observes graph only, with independent randomization allowed'),
                (a[0] > 0 and b[0] > 0, 'All within/between parameters must be strictly positive'),
                (a[0]+2*b[0] > 3, 'The entire parameter rectangle must have mean degree greater than one'),
                (max(a[1], b[1]) <= n, 'Edge probabilities a/n and b/n must not exceed one at the supplied finite n')]:
            if not condition:
                reasons.append(reason)
        result = dict(status='declared_scope_matched' if not reasons else 'outside_declared_scope',
                      scope_reasons=reasons, model_fit_verified=False,
                      source_proof_verified=False, finite_recovery_certified=False,
                      scope='Declared symmetric three-community sparse SBM with fixed connectivity parameters; no model fitting or theorem verification')
        if not reasons:
            envelope = margin_envelope(a, b)
            low, high = Fraction(envelope['minimum']), Fraction(envelope['maximum'])
            regime = 'strictly_above_for_entire_rectangle' if low > 0 else 'at_or_below_for_entire_rectangle' if high <= 0 else 'rectangle_spans_threshold'
            result.update(margin_envelope=envelope, source_conditional_regime=regime,
                          mean_degree_interval=[str((a[i]+2*b[i])/3) for i in range(2)],
                          interval_interpretation='All combinations in the supplied rectangle; no statistical coverage probability or parameter independence inferred',
                          theorem_interpretation='Source-conditional asymptotic weak-recovery statement, not a finite graph verdict or guarantee for a selected algorithm')
            if a[0] == a[1] and b[0] == b[1]:
                d = (a[0]+2*b[0])/3
                lam = (a[0]-b[0])/(a[0]+2*b[0])
                result['point_parameters'] = dict(d=str(d), lambda_value=str(lam), KS_ratio=str(d*lam**2))
        if 'evaluation' in doc:
            ev = doc['evaluation']
            if type(ev) is not dict or set(ev) != {'truth','predicted'} or q != 3:
                raise Invalid('evaluation needs exactly truth/predicted arrays for three labels')
            result['finite_score'] = score(ev['truth'], ev['predicted'])
            if result['finite_score']['count'] != n:
                raise Invalid('evaluation length must match n')
            result['exact_uniform_baseline'] = exact_uniform_label_baseline(ev['truth'])
        return result
    except (Invalid, ZeroDivisionError) as exc:
        return dict(status='invalid_input', reason=str(exc), finite_recovery_certified=False)


def unique_object(pairs):
    obj = {}
    for key, value in pairs:
        if key in obj:
            raise Invalid('duplicate JSON key')
        obj[key] = value
    return obj


def decode(raw):
    if len(raw) > MAX_BYTES:
        raise Invalid('input exceeds one MiB')
    return json.loads(raw.decode('utf-8'), object_pairs_hook=unique_object,
                      parse_constant=lambda _: (_ for _ in ()).throw(Invalid('nonfinite JSON value')))


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path)
    parser.add_argument('--output', type=Path, required=True)
    args = parser.parse_args()
    if args.output.exists():
        raise FileExistsError('Refusing to overwrite an existing output')
    with args.input.open('rb') as f:
        raw = f.read(MAX_BYTES+1)
    try:
        result = evaluate(decode(raw))
    except (ValueError, RecursionError) as exc:
        result = dict(status='invalid_input', reason=str(exc), finite_recovery_certified=False)
    result.update(input_sha256=hashlib.sha256(raw).hexdigest(), input_bytes_read=len(raw),
                  input_truncated=len(raw)>MAX_BYTES, checker_sha256=hashlib.sha256(Path(__file__).read_bytes()).hexdigest())
    with args.output.open('x') as f:
        json.dump(result, f, indent=2); f.write('\n')
    print(json.dumps(dict(status=result['status'], output=str(args.output))))


if __name__ == '__main__':
    main()
