"""A small, standard-library-only reporting example, not a benchmark run."""
import csv
import hashlib
import itertools
import json
import math
import platform
import random
import statistics as st
import time
from pathlib import Path

ROOT = Path(__file__).resolve().parent


def validate(rows, manifest):
    expected = set(itertools.product(manifest['configuration']['seeds'],
                                    manifest['configuration']['epsilon'],
                                    manifest['methods']))
    seen = set()
    for row in rows:
        key = (row['seed'], row['epsilon'], row['method'])
        if key in seen or key not in expected:
            raise ValueError('duplicate or unexpected run/condition')
        seen.add(key)
        if (row['status'] != 'ok' or row['value'] is None or
                not math.isfinite(row['value']) or not 0 <= row['value'] <= 1):
            raise ValueError('failed, missing, nonfinite or out-of-range score')
        if (row['protocol_sha256'] != manifest['protocol_sha256'] or
                row['split'] != 'test' or row['metric'] != 'robust_accuracy' or
                row['run_id'] != f"seed{row['seed']}"):
            raise ValueError('incompatible protocol or identity')
        n, failures = row['n_examples'], row['failures']
        if (n != manifest['n_examples'] or not isinstance(failures, int) or
                not 0 <= failures <= n or abs(row['value']-(n-failures)/n) > 1e-14):
            raise ValueError('inconsistent count or denominator')
    if seen != expected:
        raise ValueError('incomplete planned grid; report missingness before aggregation')


def stats(values):
    return dict(n=len(values), mean=st.mean(values), sd=st.stdev(values))


def quantile(values, p):
    """Inverse empirical CDF (nearest rank); appropriate for discrete exact mass."""
    values = sorted(values)
    return values[max(0, math.ceil(len(values)*p)-1)]


def paired_bootstrap(a, b):
    """Exact empirical bootstrap for THIS three-pair teaching example."""
    if len(a) != 3 or len(b) != 3:
        raise ValueError('enumeration demo requires exactly three aligned pairs')
    delta = [x-y for x, y in zip(a, b)]
    samples = [st.mean(delta[i] for i in indices)
               for indices in itertools.product(range(3), repeat=3)]
    rng = random.Random(80)
    mc = [st.mean(delta[rng.randrange(3)] for _ in range(3)) for _ in range(10000)]
    return dict(**stats(delta), per_seed=delta,
                percentile95=[quantile(samples, .025), quantile(samples, .975)],
                exact_resamples=27, mc_seed=80, mc_reps=len(mc),
                mc_percentile95=[quantile(mc, .025), quantile(mc, .975)],
                warning='Nominal percentile interval only; n=3, coverage not established'), samples


def build(rows, manifest):
    validate(rows, manifest)
    seeds = manifest['configuration']['seeds']
    lookup = {(r['seed'], r['epsilon'], r['method']): r['value'] for r in rows}
    groups = []
    for eps in manifest['configuration']['epsilon']:
        for method in manifest['methods']:
            values = [lookup[s, eps, method] for s in seeds]
            groups.append(dict(epsilon=eps, method=method, **stats(values), per_seed=values))
    a = [lookup[s, .2, 'pgd2'] for s in seeds]
    b = [lookup[s, .2, 'pgd20'] for s in seeds]
    paired, samples = paired_bootstrap(a, b)
    return dict(groups=groups, paired_pgd2_minus_pgd20=paired), samples


def self_checks(rows, manifest):
    import copy
    rejected = []
    cases = {'duplicate': rows+[rows[0]], 'missing': rows[:-1]}
    for name, field, value in [('failed', 'status', 'failed'), ('null', 'value', None),
                               ('nan', 'value', float('nan')), ('range', 'value', 1.1),
                               ('protocol', 'protocol_sha256', 'wrong'),
                               ('count', 'failures', -1), ('denominator', 'n_examples', 1)]:
        changed = copy.deepcopy(rows)
        changed[0][field] = value
        cases[name] = changed
    for name, case in cases.items():
        try:
            validate(case, manifest)
        except ValueError:
            rejected.append(name)
        else:
            raise AssertionError(f'accepted invalid fixture: {name}')
    # Pairing preserves a constant per-run difference despite varied raw scores.
    paired, _ = paired_bootstrap([.2, .5, .8], [.1, .4, .7])
    assert all(abs(x-.1) < 1e-14 for x in paired['percentile95'])
    identical, _ = paired_bootstrap([.2, .5, .8], [.2, .5, .8])
    assert identical['percentile95'] == [0., 0.]
    baseline, _ = build(rows, manifest)
    shuffled = list(rows)
    random.Random(80).shuffle(shuffled)
    assert build(shuffled, manifest)[0] == baseline
    return dict(rejected_fixtures=rejected, pairing_constant=True,
                identical_methods=True, order_invariant=True)


def main():
    started = time.perf_counter()
    manifest = json.loads((ROOT/'input/manifest.json').read_text())
    source = (ROOT/'input/source_summary.json').read_bytes()
    assert hashlib.sha256(source).hexdigest() == manifest['source_sha256']
    rows = [json.loads(line) for line in (ROOT/'input/runs.jsonl').read_text().splitlines()]
    result, samples = build(rows, manifest)
    checks = self_checks(rows, manifest)
    out = ROOT/'results'
    out.mkdir(exist_ok=True)
    (out/'report.json').write_text(json.dumps(result, indent=2, allow_nan=False)+'\n')
    (out/'bootstrap_exact.json').write_text(json.dumps(samples, indent=2)+'\n')
    with (out/'table.csv').open('w', newline='') as f:
        writer = csv.DictWriter(f, fieldnames=['epsilon', 'method', 'n', 'mean', 'sd'])
        writer.writeheader()
        writer.writerows({k: r[k] for k in writer.fieldnames} for r in result['groups'])
    lines = ['| Method | n runs | Mean RA (%) | Seed SD (pp) |',
             '|---|---:|---:|---:|']
    for r in result['groups']:
        if r['epsilon'] == .2:
            lines.append(f"| {r['method']} | {r['n']} | {100*r['mean']:.3f} | {100*r['sd']:.3f} |")
    (out/'table.md').write_text('\n'.join(lines)+'\n')
    log = dict(python=platform.python_version(), device='CPU', dependencies='standard library',
               source_sha256=manifest['source_sha256'], rows=len(rows),
               score_shape=[3, 5, 4], checks=checks,
               elapsed_seconds=time.perf_counter()-started,
               paired=result['paired_pgd2_minus_pgd20'])
    (out/'verification.json').write_text(json.dumps(log, indent=2)+'\n')
    print(json.dumps(log, indent=2))


if __name__ == '__main__':
    main()
