#!/usr/bin/env python3
"""Finite e+e- thrust bin: Python 3.10+, standard library only.

Run: python thrust-bin-calculation.py
Or:  python thrust-bin-calculation.py --tau-a 0.12 --tau-b 0.16

Born-normalized, massless SU(3), tree-level three-parton rate. The separate
matched result is a fixed-coupling logarithmic illustration, not full NLL
resummation or a hadron-level precision prediction. No external data or random
samples. Gauss-Legendre integration of the energy-fraction density is compared
with Simpson integration of its analytic thrust projection.

Sources: Heinrich (2025), CERN-2025-009, pp.41,49, Eqs.(98)-(99),(119)-(121);
Becher and Schwartz, arXiv:0803.0342v2, pp.1-2, Eqs.(2)-(3).
"""
import argparse
import json
import math

CF = 4.0 / 3.0


def gauss_rule(n):
    """Roots and weights of P_n, found by Newton iteration and its recurrence."""
    pairs = []
    for i in range(1, n + 1):
        z = math.cos(math.pi * (i - 0.25) / (n + 0.5))
        for _ in range(50):
            p0, p1 = 1.0, z
            for k in range(2, n + 1):
                p0, p1 = p1, ((2*k-1)*z*p1 - (k-1)*p0) / k
            derivative = n * (z*p1 - p0) / (z*z - 1.0)
            delta = p1 / derivative
            z -= delta
            if abs(delta) < 2e-15:
                break
        else:
            raise ArithmeticError("Legendre root did not converge")
        p0, p1 = 1.0, z
        for k in range(2, n + 1):
            p0, p1 = p1, ((2*k-1)*z*p1 - (k-1)*p0) / k
        derivative = n * (z*p1 - p0) / (z*z - 1.0)
        pairs.append((z, 2.0 / ((1.0-z*z)*derivative*derivative)))
    return pairs


def integrate_gauss(function, a, b, rule):
    half, mid = (b-a)/2.0, (a+b)/2.0
    return half * math.fsum(w * function(mid + half*z) for z, w in rule)


def density(x1, x2):
    return (x1*x1 + x2*x2) / ((1.0-x1)*(1.0-x2))


def projected_kernel(thrust):
    t = thrust
    return (2*(3*t*t-3*t+2)/(t*(1-t)) * math.log((2*t-1)/(1-t))
            - 3*(3*t-2)*(2-t)/(1-t))


def sectors(tau_a, tau_b, n):
    rule = gauss_rule(n)
    def sector(gluon):
        def outer(t):
            inner = (lambda x: density(x, 2-t-x)) if gluon else (lambda x: density(t, x))
            return integrate_gauss(inner, 2-2*t, t, rule)
        return integrate_gauss(outer, 1-tau_b, 1-tau_a, rule)
    return {"one_quark_maximum": sector(False), "gluon_maximum": sector(True)}


def simpson(function, a, b, intervals):
    assert intervals % 2 == 0
    h = (b-a)/intervals
    return h/3 * math.fsum(
        (1 if i in (0, intervals) else 4 if i % 2 else 2) * function(a+i*h)
        for i in range(intervals+1))


def log_coefficient(tau):
    big_l = math.log(1/tau)
    return -big_l*big_l + 1.5*big_l


def matching(alpha, tau_a, tau_b, coefficient):
    a = CF*alpha/math.pi
    ca, cb = log_coefficient(tau_a), log_coefficient(tau_b)
    # expm1 avoids cancellation for the small-coupling expansion test.
    logarithmic = math.exp(a*ca) * math.expm1(a*(cb-ca))
    expanded = a*(cb-ca)
    fixed_order = CF*alpha/(2*math.pi) * coefficient
    return {"fixed_order": fixed_order, "logarithmic": logarithmic,
            "logarithmic_first_order": expanded,
            "matched_illustration": logarithmic + fixed_order - expanded}


def calculate(tau_a=0.08, tau_b=0.12, alpha=0.12, q=100.0, nf=5):
    if not (0 < tau_a < tau_b < math.exp(-1.5)):
        raise ValueError("Use 0 < tau_a < tau_b < exp(-3/2) for this logarithmic illustration")
    if not (0 < alpha < 0.5 and q > 0 and 0 <= nf <= 6):
        raise ValueError("Require 0 < alpha_s < 0.5, Q > 0 and integer 0 <= nf <= 6")
    convergence = []
    for n in (16, 32, 64):
        values = sectors(tau_a, tau_b, n)
        coefficient = 2*values["one_quark_maximum"] + values["gluon_maximum"]
        convergence.append({"nodes_per_dimension": n, "coefficient": coefficient})
    independent = simpson(lambda tau: projected_kernel(1-tau), tau_a, tau_b, 4096)
    tolerance = 2e-10 * max(1, abs(coefficient))
    assert abs(coefficient - independent) < tolerance
    assert abs(convergence[-1]["coefficient"] - convergence[-2]["coefficient"]) < tolerance
    assert values["one_quark_maximum"] > 0 and values["gluon_maximum"] > 0
    central = matching(alpha, tau_a, tau_b, coefficient)
    beta0 = 11 - 2*nf/3
    variations = []
    for ratio in (0.5, 1.0, 2.0):
        alpha_mu = alpha / (1 + beta0*alpha/(2*math.pi)*math.log(ratio))
        variations.append({"mu_R_over_Q": ratio, "alpha_s": alpha_mu,
                           **matching(alpha_mu, tau_a, tau_b, coefficient)})

    # Direct enumeration of all hemisphere sign choices supplies an independent
    # thrust maximization for three momenta closing into a triangle.
    thrust_errors = []
    for x1, x2 in ((0.8, 0.7), (0.6, 0.8), (2/3, 2/3), (0.91, 0.83)):
        x3 = 2-x1-x2
        cosine = (x3*x3-x1*x1-x2*x2)/(2*x1*x2)
        p1, p2 = (x1, 0.0), (x2*cosine, x2*math.sqrt(1-cosine*cosine))
        momenta = [p1, p2, (-p1[0]-p2[0], -p1[1]-p2[1])]
        maximum = max(math.hypot(*[sum((1 if mask & (1 << i) else -1)*momenta[i][j]
                                          for i in range(3)) for j in range(2)])
                      for mask in range(8))/2
        thrust_errors.append(abs(maximum - max(x1, x2, x3)))
    assert max(thrust_errors) < 2e-14
    assert abs(projected_kernel(2/3)) < 1e-13
    # The integral of a quadratic checks the quadrature weights and mapping.
    assert abs(integrate_gauss(lambda x: x*x, -2, 3, gauss_rule(16)) - 35/3) < 1e-12
    # Additive matching must reproduce all of the O(alpha_s) coefficient.
    alpha_small = 1e-6
    def remainder(coupling):
        result = matching(coupling, tau_a, tau_b, coefficient)
        return result["matched_illustration"] - result["fixed_order"]
    remainder_ratio = remainder(alpha_small) / remainder(alpha_small/2)
    assert abs(remainder_ratio - 4) < 2e-5
    # Negative controls alter physical counting or matching, not tolerances.
    controls = {
        "omit_gluon_sector_absolute_error": values["gluon_maximum"],
        "omit_antiquark_sector_absolute_error": values["one_quark_maximum"],
        "double_count_matching_absolute_error": central["logarithmic_first_order"],
        "use_T_in_place_of_tau_absolute_error": abs(
            CF*alpha/math.pi*(log_coefficient(1-tau_b)-log_coefficient(1-tau_a))
            - central["logarithmic_first_order"]),
    }
    assert all(value > 1e-5 for value in controls.values())
    return {
        "format_version": 1,
        "scope": "Born-normalized massless tree-level bin plus a fixed-coupling logarithmic matching illustration; no full NLL, hadronization, or detector model",
        "inputs": {"Q_GeV": q, "alpha_s_Q": alpha, "n_f": nf, "C_F": CF,
                   "tau_a": tau_a, "tau_b": tau_b, "normalization": "sigma_0"},
        "central": central, "phase_space_sectors": values,
        "quadrature_convergence": convergence,
        "independent_projected_coefficient": independent,
        "absolute_quadrature_difference": abs(coefficient-independent),
        "renormalization_scale_diagnostic": variations,
        "natural_scales_GeV": {"hard": q, "jet_range": [q*math.sqrt(tau_a), q*math.sqrt(tau_b)],
                               "soft_range": [q*tau_a, q*tau_b]},
        "checks": {"maximum_thrust_error": max(thrust_errors),
                   "matching_small_coupling_remainder_ratio": remainder_ratio,
                   "negative_controls": controls},
    }


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--tau-a", type=float, default=0.08)
    parser.add_argument("--tau-b", type=float, default=0.12)
    parser.add_argument("--alpha-s", type=float, default=0.12)
    parser.add_argument("--Q", type=float, default=100.0)
    parser.add_argument("--nf", type=int, default=5)
    args = parser.parse_args()
    print(json.dumps(calculate(args.tau_a, args.tau_b, args.alpha_s, args.Q, args.nf), indent=2))
