"""Reproduce the principal analyses for the Grade 6 JrMAI-B study.

Run:
    python reproduce_analysis.py

Required packages: pandas, numpy, scipy, statsmodels, openpyxl.
The script expects Supplementary_Data.xlsx in the same directory.
"""

from pathlib import Path
import math
import numpy as np
import pandas as pd
from scipy import stats

HERE = Path(__file__).resolve().parent
DATA = HERE / "Supplementary_Data.xlsx"
OUT = HERE / "statistical_output.txt"

KOC = [1, 2, 3, 4, 5, 12, 13, 14, 16]
ROC = [6, 7, 8, 9, 10, 11, 15, 17, 18]
CLUSTERS = {
    "declarative": [1, 4, 12],
    "conditional": [2, 5, 13, 14],
    "procedural": [3, 16],
    "information_management": [6, 11],
    "evaluation": [7, 17],
    "monitoring": [8, 10, 15],
    "planning": [9, 18],
}


def alpha(frame):
    values = frame.to_numpy(float)
    k = values.shape[1]
    return k / (k - 1) * (1 - values.var(axis=0, ddof=1).sum() /
                           values.sum(axis=1).var(ddof=1))


def fisher_ci(r, n, level=.95):
    z = np.arctanh(r)
    se = 1 / math.sqrt(n - 3)
    crit = stats.norm.ppf(1 - (1 - level) / 2)
    return tuple(np.tanh([z - crit * se, z + crit * se]))


def welch_anova(samples):
    """Welch one-way ANOVA with Satterthwaite denominator degrees of freedom."""
    samples = [np.asarray(x, dtype=float) for x in samples]
    n = np.array([len(x) for x in samples], dtype=float)
    means = np.array([x.mean() for x in samples])
    variances = np.array([x.var(ddof=1) for x in samples])
    weights = n / variances
    k = len(samples)
    weighted_mean = np.sum(weights * means) / np.sum(weights)
    numerator = np.sum(weights * (means - weighted_mean) ** 2) / (k - 1)
    term = np.sum((1 - weights / np.sum(weights)) ** 2 / (n - 1))
    correction = 1 + (2 * (k - 2) / (k**2 - 1)) * term
    f_value = numerator / correction
    df1 = k - 1
    df2 = (k**2 - 1) / (3 * term)
    p_value = stats.f.sf(f_value, df1, df2)
    return f_value, df1, df2, p_value


def main():
    d = pd.read_excel(DATA, sheet_name="Raw_Data")
    q = {i: f"jmai_{i:02d}" for i in range(1, 19)}
    d["science_mean"] = d[["science_q1", "science_q2", "science_q3"]].mean(axis=1)
    d["knowledge_mean"] = d[[q[i] for i in KOC]].mean(axis=1)
    d["regulation_mean"] = d[[q[i] for i in ROC]].mean(axis=1)
    d["overall_jmai_mean"] = d[[q[i] for i in range(1, 19)]].mean(axis=1)
    d["achievement_group"] = pd.cut(
        d["science_mean"], [-np.inf, 75, 80, 85, 90, np.inf],
        right=False,
        labels=["Did not meet expectations", "Fairly satisfactory",
                "Satisfactory", "Very satisfactory", "Outstanding"],
    )
    for name, items in CLUSTERS.items():
        d[name] = d[[q[i] for i in items]].mean(axis=1)

    lines = []
    lines.append(f"N = {len(d)}; complete rows = {d.notna().all(axis=1).sum()}")
    lines.append("\nDESCRIPTIVE STATISTICS")
    for col in ["science_mean", "overall_jmai_mean", "knowledge_mean", "regulation_mean"]:
        x = d[col]
        ci = stats.t.interval(.95, len(x)-1, loc=x.mean(), scale=stats.sem(x))
        lines.append(f"{col}: M={x.mean():.3f}, SD={x.std(ddof=1):.3f}, "
                     f"range=[{x.min():.3f}, {x.max():.3f}], 95% CI={ci}")

    lines.append("\nINTERNAL CONSISTENCY")
    lines.append(f"overall alpha = {alpha(d[[q[i] for i in range(1,19)]]):.3f}")
    lines.append(f"knowledge alpha = {alpha(d[[q[i] for i in KOC]]):.3f}")
    lines.append(f"regulation alpha = {alpha(d[[q[i] for i in ROC]]):.3f}")

    lines.append("\nACHIEVEMENT CATEGORIES")
    lines.append(d["achievement_group"].value_counts(sort=False).to_string())

    lines.append("\nCORRELATIONS WITH EXACT SCIENCE MEAN")
    for col in ["overall_jmai_mean", "knowledge_mean", "regulation_mean", *CLUSTERS]:
        r, p = stats.pearsonr(d["science_mean"], d[col])
        rho, ps = stats.spearmanr(d["science_mean"], d[col])
        lines.append(f"{col}: Pearson r={r:.3f}, 95% CI={fisher_ci(r,len(d))}, "
                     f"p={p:.8g}, r2={r*r:.3f}; Spearman rho={rho:.3f}, p={ps:.8g}")

    lines.append("\nWELCH ANOVA")
    observed = d[d["achievement_group"] != "Did not meet expectations"]
    for col in ["overall_jmai_mean", "knowledge_mean", "regulation_mean"]:
        samples = [g[col].to_numpy() for _, g in observed.groupby(
            "achievement_group", observed=True, sort=False)]
        f_value, df1, df2, p_value = welch_anova(samples)
        lines.append(f"{col}: F={f_value:.3f}, df=({df1:.0f}, {df2:.2f}), p={p_value:.8g}")

    lines.append("\nGROUP DESCRIPTIVES")
    lines.append(observed.groupby("achievement_group", observed=True)[
        ["overall_jmai_mean", "knowledge_mean", "regulation_mean"]
    ].agg(["count", "mean", "std"]).to_string())

    lines.append("\nSCORING SENSITIVITY FOR ITEM 16")
    reversed_16 = d[[q[i] for i in KOC]].copy()
    reversed_16[q[16]] = 6 - reversed_16[q[16]]
    lines.append(f"knowledge alpha, primary positive key = "
                 f"{alpha(d[[q[i] for i in KOC]]):.3f}")
    lines.append(f"knowledge alpha, item 16 reversed = {alpha(reversed_16):.3f}")

    OUT.write_text("\n".join(lines), encoding="utf-8")
    print("\n".join(lines))


if __name__ == "__main__":
    main()
