#!/usr/bin/env python3
"""Meaningful independent numerical and analytic checks for the GAD prototype."""
import json
import math
import random
from pathlib import Path
from gad_capacity import binary_entropy, estimate, scalar_holevo

ROOT = Path(__file__).resolve().parents[2]


def state_entropy(bloch):
    norm = math.sqrt(sum(x*x for x in bloch))
    if norm > 1 + 1e-12:
        raise ArithmeticError('Nonphysical output Bloch vector')
    return binary_entropy((1 - min(1.0, norm)) / 2)


def random_ensemble_holevo(gamma, nu, rng):
    count = rng.randrange(2, 7)
    weights = [rng.random() for _ in range(count)]
    total = sum(weights)
    weights = [w / total for w in weights]
    outputs = []
    for _ in range(count):
        z = rng.uniform(-1, 1)
        angle = rng.uniform(0, 2 * math.pi)
        xy = math.sqrt(1-z*z)
        outputs.append((math.sqrt(1-gamma)*xy*math.cos(angle),
                        math.sqrt(1-gamma)*xy*math.sin(angle),
                        (1-gamma)*z+gamma*(1-2*nu)))
    mean = [sum(w*v[j] for w,v in zip(weights, outputs)) for j in range(3)]
    return state_entropy(mean)-sum(w*state_entropy(v) for w,v in zip(weights, outputs))


def main():
    checks = []
    def check(name, error, tolerance):
        if error > tolerance:
            raise AssertionError((name, error, tolerance))
        checks.append(dict(name=name, absolute_error=error, tolerance=tolerance, passed=True))

    for nu in (0.0, 0.17, 0.5, 1.0):
        check(f'identity_nu_{nu}', abs(estimate(0, nu)['best_observed_holevo_bits']-1), 1e-12)
        check(f'replacement_nu_{nu}', abs(estimate(1, nu)['best_observed_holevo_bits']), 1e-12)
    for gamma in (0.03, 0.2, 0.6, 0.95):
        expected = 1-binary_entropy((1-math.sqrt(1-gamma))/2)
        check(f'unital_gamma_{gamma}', abs(estimate(gamma,0.5)['best_observed_holevo_bits']-expected), 1e-11)
    for gamma,nu in ((0.2,0.1),(0.55,0.3),(0.9,0.03)):
        a,b=estimate(gamma,nu),estimate(gamma,1-nu)
        check(f'thermal_symmetry_{gamma}_{nu}', abs(a['best_observed_holevo_bits']-b['best_observed_holevo_bits']), 1e-11)
        for p in (0.0,0.23,0.5,0.79,1.0):
            t=(1-gamma)*p+gamma*nu
            off=math.sqrt((1-gamma)*p*(1-p))
            determinant=t*(1-t)-off*off
            direct_entropy=binary_entropy((1-math.sqrt(max(0,1-4*determinant)))/2)
            direct_value=binary_entropy(t)-direct_entropy
            check(f'direct_matrix_{gamma}_{nu}_{p}', abs(direct_value-scalar_holevo(gamma,nu,p)), 1e-11)
    rng=random.Random(276)
    numerical_ensemble_checks=[]
    for gamma,nu in ((0.2,0.1),(0.55,0.3),(0.9,0.03)):
        reference=estimate(gamma,nu)['best_observed_holevo_bits']
        sample_best=max(random_ensemble_holevo(gamma,nu,rng) for _ in range(500))
        if sample_best > reference+1e-10:
            raise AssertionError('Random ensemble exceeds scalar benchmark')
        numerical_ensemble_checks.append(dict(gamma=gamma,nu=nu,ensembles_sampled=500,
                                               best_random_holevo_bits=sample_best,scalar_benchmark_bits=reference))
    for gamma,nu in ((0.02,0.08),(0.5,0.27),(0.97,0.9)):
        coarse=estimate(gamma,nu,64)
        fine=estimate(gamma,nu,16384)
        check(f'grid_allowance_{gamma}_{nu}', max(0,fine['best_observed_holevo_bits']-coarse['estimated_model_maximum_upper_with_grid_allowance_bits']), 1e-12)
    result=dict(checks=checks,random_ensemble_checks=numerical_ensemble_checks,
                theorem_verification='not_run', certified_numeric_error_bound='not_established',
                caveat='Random ensembles and finite grid comparisons are diagnostics, not an optimality proof.')
    out=ROOT/'research/snapshots/2026-10-08-baseline/gad-validation-2026-10-09-v1.json'
    with out.open('x') as stream:
        json.dump(result,stream,indent=2);stream.write('\n')
    print(json.dumps(dict(analytic_and_matrix_checks=len(checks),random_ensembles=1500,output=str(out))))


if __name__=='__main__':
    main()
