"""Audit saved results with arithmetic independent of the generation loop."""
import argparse
import json
from pathlib import Path

import numpy as np

from minimal_decode import corrected


def main(root):
    data = json.loads((root / 'results.json').read_text())
    assert data['kv']['bytes'] == 2 * 2 * 2 * 29 * 8 * 8
    assert data['kv']['max_logit_error'] < 1e-10
    assert len(data['greedy_checks']) == 9
    for row in data['greedy_checks']:
        assert row['greedy_equal']
        assert row['rounds'] + row['accepted'] == 48
        assert row['accepted'] <= row['examined'] <= row['proposed']
        assert row['full_rounds'] + row['rejected_rounds'] == row['rounds']
    for run in data['sampling_cache_traces']:
        position, draft_length = 8, 7
        for row in run['rounds']:
            assert row['prefix_length'] == position
            assert row['accepted'] <= row['proposed'] <= 4
            if row['proposed']:
                draft_length = position + row['proposed'] - 1
            position += row['accepted'] + 1
            assert row['target_cache_length'] == position - 1
            draft_length = min(draft_length, position - 1)
            assert row['draft_cache_length'] == draft_length
        assert position == 56
    w = data['workload']
    assert w['full_prefix']['target_layer_tokens'] == 2 * sum(range(8, 56))
    assert w['full_prefix']['target_score_cells'] == 4 * sum(n*n for n in range(8, 56))
    assert w['kv_cache']['target_layer_tokens'] == 2 * (8 + 47)
    assert w['kv_cache']['target_score_cells'] == 4 * (8*8 + sum(range(9, 56)))
    greedy = next(x for x in data['greedy_checks'] if x['seed'] == 69 and x['gamma'] == 4)
    assert w['speculative_g4']['target_calls'] == greedy['rounds'] + 1
    assert w['speculative_g4']['draft_calls'] == greedy['proposed'] + 1
    assert w['speculative_g4']['target_layer_tokens'] == 2 * (7 + greedy['rounds'] + greedy['proposed'])
    for record in data['timing'].values():
        assert len(record['ms']) == 7 and min(record['ms']) > 0
        assert sorted(record['ms'])[3] == record['median_ms']
    # Use implementation's residual with a separately enumerated acceptance law.
    for p, q in [([.1,.6,.3],[.7,.2,.1]), ([0,1,0],[1,0,0]), ([.4,.6,0],[0,.5,.5])]:
        p, q = np.array(p, dtype=float), np.array(q, dtype=float)
        ratio = np.divide(p, q, out=np.zeros_like(p), where=q > 0)
        mass = q * np.minimum(1, ratio)
        distribution = mass + (1 - mass.sum()) * corrected(p, q)
        np.testing.assert_allclose(distribution, p, atol=1e-14)
    check = data['distribution_check']
    p, q = np.array(check['p']), np.array(check['q'])
    for label, tv in [('frequency','tv'), ('wrong_fallback_frequency','wrong_fallback_tv')]:
        observed = np.array(check[label])
        assert abs(observed.sum() - 1) < 1e-12
        assert abs(np.abs(observed - p).sum()/2 - check[tv]) < 1e-12
    wrong = np.minimum(p,q) + (1-np.minimum(p,q).sum())*p
    np.testing.assert_allclose(wrong, [.16,.56,.28], atol=1e-14)
    assert data['identity_draft']['accepted'] == data['identity_draft']['proposed']
    print('PASS: saved trace invariants, workload accounting, medians, exact residual law, negative control')
    print('This audit does not claim paper benchmark reproduction or prove every possible generation path.')


if __name__ == '__main__':
    parser = argparse.ArgumentParser()
    parser.add_argument('--out', type=Path, default=Path(__file__).resolve().parent)
    main(parser.parse_args().out)
