"""Reproduce a preregistered aggregate mortality decomposition, offline.

All point estimates use exact rational arithmetic. Conditional uncertainty
assumes independent Poisson death counts and fixed source populations.
"""
from fractions import Fraction as Q
from itertools import permutations
from pathlib import Path
from statistics import NormalDist
import hashlib
import json
import math
import xml.etree.ElementTree as ET

ROOT = Path(__file__).resolve().parents[1]
AGES = ["< 1 year", "1-4 years", "5-14 years", "15-24 years", "25-34 years",
        "35-44 years", "45-54 years", "55-64 years", "65-74 years",
        "75-84 years", "85+ years"]
FACTORS = ("population", "age_structure", "rates")

def integer(value):
    if isinstance(value, str):
        value = value.strip()
    if value is None or not value.replace(",", "").isdigit():
        raise ValueError("Unavailable or suppressed count: " + str(value))
    return int(value.replace(",", ""))

def parse_response(path):
    xml = ET.parse(path)
    rows, totals, unknown = {}, {}, {}
    year = None
    for row in xml.findall(".//data-table/r"):
        cells = list(row)
        if cells[0].get("r"):
            year = integer(cells.pop(0).get("l"))
        if cells[0].get("c") == "1":
            totals[year] = {"deaths": integer(cells[1].get("dt")),
                            "population": integer(cells[2].get("dt"))}
        elif cells[0].get("c") == "2":
            continue
        else:
            label = cells[0].get("l")
            if label == "Not Stated":
                value = cells[1].get("v")
                unknown[year] = value if value == "Suppressed" else integer(value)
                continue
            if label not in AGES:
                raise ValueError("Unexpected age label: " + str(label))
            key = (year, label)
            if key in rows:
                raise ValueError("Duplicate source cell")
            deaths, population = integer(cells[1].get("v")), integer(cells[2].get("v"))
            if deaths < 10 or population <= 0:
                raise ValueError("Cell fails disclosure or denominator rule")
            rows[key] = {"age": label, "deaths": deaths, "population": population}
    for year in range(2018, 2025):
        selected = [rows[(year, age)] for age in AGES]
        if year in totals:
            assert sum(r["population"] for r in selected) == totals[year]["population"]
            # WONDER may omit totals when a not-stated-age cell is suppressed.
            assert sum(r["deaths"] for r in selected) <= totals[year]["deaths"]
    return rows, totals, unknown

def factors(rows):
    population = sum(row["population"] for row in rows)
    shares = [Q(row["population"], population) for row in rows]
    rates = [Q(row["deaths"], row["population"]) for row in rows]
    return population, shares, rates

def expected(n, s, r):
    return n * sum(a*b for a, b in zip(s, r))

def permutation_decomposition(start, end):
    paths = []
    for order in permutations(range(3)):
        current = list(start)
        before = expected(*current)
        contribution = [Q(0), Q(0), Q(0)]
        for factor in order:
            current[factor] = end[factor]
            after = expected(*current)
            contribution[factor] = after - before
            before = after
        paths.append((order, contribution))
    averaged = [sum(p[1][i] for p in paths)/6 for i in range(3)]
    return averaged, paths

def expansion_decomposition(start, end):
    n, s, r = start
    dn = end[0] - n
    ds = [b-a for a, b in zip(s, end[1])]
    dr = [b-a for a, b in zip(r, end[2])]
    dot = lambda a, b: sum(x*y for x, y in zip(a, b))
    n_main, s_main, r_main = dn*dot(s, r), n*dot(ds, r), n*dot(s, dr)
    ns, nr, sr, nsr = dn*dot(ds, r), dn*dot(s, dr), n*dot(ds, dr), dn*dot(ds, dr)
    return [n_main+(ns+nr)/2+nsr/3,
            s_main+(ns+sr)/2+nsr/3,
            r_main+(nr+sr)/2+nsr/3]

def conditional_errors(start, end, old, new):
    # Exact linear coefficients, verified independently by unit-count perturbations.
    n0, s0, _ = start
    n1, s1, _ = end
    dn = n1-n0
    variances = [Q(0), Q(0), Q(0)]
    combined_variance = Q(0)
    base, _ = permutation_decomposition(start, end)
    for i in range(len(old)):
        ds = s1[i]-s0[i]
        rate_weight = Q(n0,3)*s0[i] + Q(n1,6)*s0[i] + Q(n0,6)*s1[i] + Q(n1,3)*s1[i]
        coefficients = [
            [Q(dn,6)*(2*s0[i]+s1[i]), Q(dn,6)*(s0[i]+2*s1[i])],
            [Q(2*n0+n1,6)*ds, Q(n0+2*n1,6)*ds],
            [-rate_weight, rate_weight],
        ]
        for t, original in [(0,old),(1,new)]:
            weight = (coefficients[0][t]+coefficients[1][t])/original[i]["population"]
            combined_variance += weight*weight*original[i]["deaths"]
        for f, (a, b) in enumerate(coefficients):
            a, b = a/old[i]["population"], b/new[i]["population"]
            variances[f] += a*a*old[i]["deaths"] + b*b*new[i]["deaths"]
            # Recompute all paths after one endpoint count increases by one.
            # This checks the variance weights against the separate path algorithm.
            for endpoint, coefficient, original in [(0,a,old),(1,b,new)]:
                changed = [dict(row) for row in original]
                changed[i]["deaths"] += 1
                pair = [start, end]
                pair[endpoint] = factors(changed)
                perturbed, _ = permutation_decomposition(*pair)
                assert perturbed[f]-base[f] == coefficient
    return [math.sqrt(float(v)) for v in variances], math.sqrt(float(combined_variance))

def aggregate(rows, groups):
    return [{"age": label,
             "deaths": sum(rows[i]["deaths"] for i in indices),
             "population": sum(rows[i]["population"] for i in indices)}
            for label, indices in groups]

def compare(old, new, label, y0, y1):
    start, end = factors(old), factors(new)
    values, paths = permutation_decomposition(start, end)
    independent = expansion_decomposition(start, end)
    reverse, _ = permutation_decomposition(end, start)
    observed = sum(r["deaths"] for r in new)-sum(r["deaths"] for r in old)
    assert values == independent
    assert reverse == [-v for v in values]
    assert sum(values) == observed
    for _, path in paths:
        assert sum(path) == observed
    se, combined_se = conditional_errors(start, end, old, new)
    z = NormalDist().inv_cdf(0.975)
    baseline = sum(row["deaths"] for row in old)
    components = {name: {
        "deaths": round(float(value), 1),
        "baseline_percent": round(100*float(value)/baseline, 3),
        "conditional_poisson_se_deaths": round(error, 1),
        "conditional_poisson_ci95_low_deaths": round(float(value)-z*error, 1),
        "conditional_poisson_ci95_high_deaths": round(float(value)+z*error, 1),
        "all_orders_min_deaths": round(float(min(p[1][i] for p in paths)), 1),
        "all_orders_max_deaths": round(float(max(p[1][i] for p in paths)), 1),
    } for i, (name, value, error) in enumerate(zip(FACTORS, values, se))}
    return {"label": label, "baseline_year": y0, "endpoint_year": y1,
            "baseline_known_age_deaths": baseline,
            "endpoint_known_age_deaths": sum(row["deaths"] for row in new),
            "baseline_population": start[0], "endpoint_population": end[0],
            "population_percent_change": round(100*(end[0]/start[0]-1), 3),
            "known_age_death_difference": observed,
            "death_percent_change": round(100*observed/baseline, 3),
            "combined_demographic_contribution_deaths":round(float(sum(values[:2])),1),
            "combined_demographic_conditional_ci95_low_deaths":round(float(sum(values[:2]))-z*combined_se,1),
            "combined_demographic_conditional_ci95_high_deaths":round(float(sum(values[:2]))+z*combined_se,1),
            "baseline_crude_per_100000": round(100*1000*baseline/start[0], 3),
            "endpoint_crude_per_100000": round(100*1000*sum(r["deaths"] for r in new)/end[0], 3),
            "components": components,
            "orders": [{"order": [FACTORS[i] for i in order],
                        "deaths": {name: round(float(v), 1) for name, v in zip(FACTORS, vals)}} for order, vals in paths],
            "exact_checks": {"sum_to_observed": True, "expansion_matches": True,
                             "reverse_negates": True, "each_order_reconciles": True,
                             "poisson_coefficients_match_count_perturbations": True}}

def main():
    sources = {}
    for sex in ["all", "female", "male"]:
        sources[sex] = parse_response(ROOT / "data" / ("response-"+sex+".xml"))
    def rows(sex, year):
        return [sources[sex][0][(year, age)] for age in AGES]
    # Independent sex totals also check that the all-sex known-age cells agree.
    for year in range(2018, 2025):
        for age in AGES:
            for value in ["population", "deaths"]:
                assert sources["all"][0][(year, age)][value] == sum(sources[s][0][(year, age)][value] for s in ["female", "male"])
    primary = compare(rows("all",2018), rows("all",2024), "All residents with known age",2018,2024)
    sensitivity = [compare(rows(s,2018),rows(s,2024),s,2018,2024) for s in ["female","male"]]
    sensitivity.append(compare(rows("all",2019),rows("all",2024),"Baseline 2019",2019,2024))
    coarse = [("Below 65",range(8)),("65-74",[8]),("75-84",[9]),("85+",[10])]
    sensitivity.append(compare(aggregate(rows("all",2018),coarse),aggregate(rows("all",2024),coarse),"Coarser ages",2018,2024))
    sensitivity.append(compare(rows("all",2018)[:-1],rows("all",2024)[:-1],"Below 85 only",2018,2024))
    annual = [compare(rows("all",y),rows("all",y+1),"Adjacent years",y,y+1) for y in range(2018,2024)]
    source_summary = [{"year":y,"all_age_deaths":sources["all"][1][y]["deaths"],
                       "known_age_deaths":sum(r["deaths"] for r in rows("all",y)),
                       "unknown_age_deaths":sources["all"][2].get(y,"No row"),
                       "population":sources["all"][1][y]["population"]} for y in range(2018,2025)]
    # Source-reported unknown-age totals are published only if >=10.
    for row in source_summary:
        if isinstance(row["unknown_age_deaths"],int) and row["unknown_age_deaths"]<10:
            row["unknown_age_deaths"]="Below disclosure threshold; not displayed"
    result = {"primary": primary, "sensitivity":sensitivity,"annual":annual,
              "rate_reporting_population":100000,
              "endpoint_age_cells":[{"age":a,
                  "baseline_deaths":sources["all"][0][(2018,a)]["deaths"],
                  "endpoint_deaths":sources["all"][0][(2024,a)]["deaths"],
                  "baseline_population":sources["all"][0][(2018,a)]["population"],
                  "endpoint_population":sources["all"][0][(2024,a)]["population"],
                  "baseline_rate_per_100000":round(100000*sources["all"][0][(2018,a)]["deaths"]/sources["all"][0][(2018,a)]["population"],3),
                  "endpoint_rate_per_100000":round(100000*sources["all"][0][(2024,a)]["deaths"]/sources["all"][0][(2024,a)]["population"],3)} for a in AGES],
              "source_summary":source_summary,
              "source_cell_count":sum(len(s[0]) for s in sources.values()),
              "source_sha256":{p.name:hashlib.sha256(p.read_bytes()).hexdigest() for p in sorted((ROOT/"data").glob("response-*.xml"))},
              "checks":{"sex_cells_sum_to_all":True,"all_sex_known_age_populations_match_source_totals":True,
                        "sex_specific_source_totals":"Omitted by WONDER; validate against the all-sex known-age cells, without reconstructing suppressed cells"},
              "uncertainty_model":"Independent Poisson counts conditional on fixed populations; analytic linear variances; normal 95% intervals"}
    out = ROOT/"results";out.mkdir(exist_ok=True)
    (out/"R1.json").write_text(json.dumps(result,ensure_ascii=False,indent=2)+"\n")
    print(json.dumps(primary,indent=2))

if __name__ == "__main__":
    main()
