"""Independent scalar audit of saved witnesses; does not import training code."""
import itertools
import json
import math
import statistics
from pathlib import Path
import numpy as np

ROOT = Path(__file__).resolve().parent
OUT = ROOT / 'results'


def main():
    summary = json.loads((OUT/'summary.json').read_text())
    errors, predictions, constraints = [], 0, 0
    for run in summary['runs']:
        with np.load(OUT/f"seed{run['seed']}.npz") as f:
            # Load once: repeated NPZ decompression would dominate this audit.
            saved = {k: f[k] for k in f.files}
        x, y, w = saved['test'].tolist(), saved['test_y'].tolist(), saved['w'].tolist()
        b = float(saved['b'])
        clean = [label*(math.fsum(a*c for a,c in zip(v,w))+b) for v,label in zip(x,y)]
        for k, budget in enumerate(run['budgets']):
            eps = budget['epsilon']
            # Coordinate-wise interval minimization, independently of vector formula.
            exact = [label*b + math.fsum(min(label*wj*(v-eps), label*wj*(v+eps))
                                        for v,wj in zip(row,w)) for row,label in zip(x,y)]
            errors.extend(abs(a-bb) for a,bb in zip(exact,saved[f'e{k}_exact_margin']))
            assert sum(v>0 for v in exact)/len(y) == budget['exact_robust_accuracy']
            union = [m>0 for m in clean]
            for name, record in budget['attacks'].items():
                candidates = saved[f'e{k}_{name}'].tolist()
                margins = []
                for original, z, label, lower in zip(x,candidates,y,exact):
                    assert all(abs(zi-xi) <= eps+1e-12 for xi,zi in zip(original,z))
                    assert all(math.isfinite(zi) for zi in z)
                    m = label*(math.fsum(a*c for a,c in zip(z,w))+b)
                    assert m >= lower-1e-10
                    margins.append(m)
                    constraints += len(z)
                clean_count = sum(m>0 for m in clean)
                survivors = [c>0 and a>0 for c,a in zip(clean,margins)]
                n = len(y)
                rebuilt = dict(clean_accuracy=clean_count/n, robust_accuracy=sum(survivors)/n,
                    conditional_asr=sum(c>0 and not s for c,s in zip(clean,survivors))/clean_count,
                    clean_correct=clean_count, failures=n-sum(survivors),
                    missed_failures=sum(s and e<=0 for s,e in zip(survivors,exact)),
                    max_margin_gap=max(a-e for a,e in zip(margins,exact)))
                for key,val in rebuilt.items():
                    errors.append(abs(val-record[key]))
                if name in ('fgsm','pgd20'):
                    errors.extend(abs(a-e) for a,e in zip(margins,exact))
                if name == 'zero_gradient_fixture':
                    assert candidates == x
                union = [a and bb for a,bb in zip(union,survivors)]
                predictions += n
            assert sum(union)/len(y) == budget['union_robust_accuracy']
        rows = [tuple(row) for split in ('train','dev','test') for row in saved[split].tolist()]
        assert len(set(rows)) == len(rows) == 3328
    for k, budget in enumerate(summary['aggregate']):
        for name, metric_set in budget['attacks'].items():
            for key, stats in metric_set.items():
                values = [r['budgets'][k]['attacks'][name][key] for r in summary['runs']]
                errors.extend([abs(statistics.mean(values)-stats['mean']),
                               abs(statistics.stdev(values)-stats['sd'])])
    # All corners of a tiny domain: zeros in weights and both labels included.
    for label in (-1,1):
        point, weights, bias, eps = [0.2,-0.3,0.8], [1.3,-0.7,0.], -.2, .4
        corners = [label*(sum((a+d)*v for a,d,v in zip(point,ds,weights))+bias)
                   for ds in itertools.product((-eps,eps),repeat=3)]
        bound = label*(sum(a*v for a,v in zip(point,weights))+bias)-eps*sum(map(abs,weights))
        errors.append(abs(min(corners)-bound))
    # Threat-domain counterexample: pixel clipping changes the unconstrained bound.
    pixel, weight, eps = .02, 1., .1
    assert pixel*weight-eps*abs(weight) == -.08
    assert max(0.,pixel-eps)*weight == 0.
    templates = [json.loads(line) for line in (ROOT/'text_templates.jsonl').read_text().splitlines()]
    assert len({r['id'] for r in templates}) == len(templates) == 6
    assert sum(r['label_preserved'] for r in templates) == 4
    assert all(r['inference_status']=='not_run' for r in templates)
    assert max(errors) < 1e-9
    report = dict(status='passed', scalar_candidate_predictions=predictions,
        coordinate_budget_checks=constraints, max_absolute_error=max(errors),
        corner_fixture_count=16, disjoint_rows_per_seed=3328,
        template_schema_rows=6, template_model_inference='not_run',
        scope='independent arithmetic/metric audit; not a neural benchmark')
    (OUT/'audit.json').write_text(json.dumps(report,indent=2)+'\n')
    print(json.dumps(report,indent=2))


if __name__ == '__main__':
    main()
