# Reviewer's independent six-order (Shapley) decomposition of D = N * sum_a s_a r_a from the
# bundle's endpoint age cells, and its sensitivity to the 2024 denominators of ages under 65
# (where the Vintage 2024 net-international-migration revision mostly falls).
import itertools, json
cells = json.load(open("bundle/results/R1.json"))["endpoint_age_cells"]
d0 = [c["baseline_deaths"] for c in cells]; d1 = [c["endpoint_deaths"] for c in cells]
n0 = [c["baseline_population"] for c in cells]; n1 = [c["endpoint_population"] for c in cells]
ages = [c["age"] for c in cells]
def factors(d, n):
    N = sum(n); return {"population": N, "age_structure": [x / N for x in n], "rates": [a / b for a, b in zip(d, n)]}
def D(f): return f["population"] * sum(s * r for s, r in zip(f["age_structure"], f["rates"]))
def shapley(d0, n0, d1, n1):
    f0, f1 = factors(d0, n0), factors(d1, n1); out = {k: 0.0 for k in f0}
    orders = list(itertools.permutations(f0))
    for o in orders:
        cur = dict(f0)
        for k in o:
            before = D(cur); cur[k] = f1[k]; out[k] += (D(cur) - before) / len(orders)
    return {k: round(v, 1) for k, v in out.items()}
print("reproduced primary:", shapley(d0, n0, d1, n1))
young = [i for i, a in enumerate(ages) if not any(t in a for t in ("65", "75", "85"))]
for pct in (1, 2, 3):
    n1p = [x * (1 - pct / 100) if i in young else x for i, x in enumerate(n1)]
    print(f"2024 under-65 population {pct}% lower:", shapley(d0, n0, d1, n1p))
