from pathlib import Path
from decimal import Decimal, ROUND_HALF_UP
import platform

import numpy as np
import pandas as pd
import scipy
from scipy.stats import kruskal, spearmanr
import scikit_posthocs as sp
import pingouin as pg
import sklearn
from sklearn.metrics import cohen_kappa_score


DATA_FILE = Path(__file__).with_name("Additional file 1.xlsx")
OUTPUT_DIR = Path(__file__).with_name("analysis_outputs")
SEARCH_DATE = pd.Timestamp("2026-05-12")

SOURCE_ORDER = [
    "Public institution",
    "Patient experience",
    "Private institution",
    "Physician",
    "Health-related media channel",
]

ANALYSIS_VARIABLES = [
    "Views",
    "Likes",
    "Comments",
    "Video age (months)",
    "Mean DISCERN",
    "Mean JAMA",
    "Mean GQS",
]


def round_half_up(value, digits=2):
    quantum = Decimal("1").scaleb(-digits)
    return Decimal(str(float(value))).quantize(quantum, rounding=ROUND_HALF_UP)


def calculate_video_age(upload_date):
    upload_date = pd.to_datetime(upload_date)
    return (
        (SEARCH_DATE.year - upload_date.dt.year) * 12
        + (SEARCH_DATE.month - upload_date.dt.month)
        + (SEARCH_DATE.day - upload_date.dt.day) / 30.0
    )


def load_data():
    df = pd.read_excel(DATA_FILE, sheet_name="S1 Data")

    required = [
        "Video ID",
        "Video URL",
        "Video title",
        "Upload source",
        "Views",
        "Likes",
        "Comments",
        "Upload date",
        "Country",
        "DISCERN rater 1",
        "JAMA rater 1",
        "GQS rater 1",
        "DISCERN rater 2",
        "JAMA rater 2",
        "GQS rater 2",
    ]
    missing_columns = [column for column in required if column not in df.columns]
    if missing_columns:
        raise ValueError(f"Missing columns: {missing_columns}")

    df = df.dropna(subset=["Video ID"]).copy()
    df["Upload date"] = pd.to_datetime(df["Upload date"])
    df["Video age (months)"] = calculate_video_age(df["Upload date"])

    df["Mean DISCERN"] = df[
        ["DISCERN rater 1", "DISCERN rater 2"]
    ].mean(axis=1)
    df["Mean JAMA"] = df[
        ["JAMA rater 1", "JAMA rater 2"]
    ].mean(axis=1)
    df["Mean GQS"] = df[
        ["GQS rater 1", "GQS rater 2"]
    ].mean(axis=1)
    df["Composite quality score"] = df[
        ["Mean DISCERN", "Mean JAMA", "Mean GQS"]
    ].sum(axis=1)

    if df[required].isna().any().any():
        raise ValueError("Missing data were detected in the analysis variables.")

    df["Upload source"] = pd.Categorical(
        df["Upload source"],
        categories=SOURCE_ORDER,
        ordered=True,
    )
    return df


def descriptive_statistics(df):
    result = df[ANALYSIS_VARIABLES].agg(
        ["mean", "std", "min", "max"]
    ).T
    result.columns = ["Mean", "SD", "Minimum", "Maximum"]
    return result


def source_summary(df):
    summary = (
        df.groupby("Upload source", observed=True)[ANALYSIS_VARIABLES]
        .agg(["count", "mean", "std"])
        .reindex(SOURCE_ORDER)
    )
    return summary


def kruskal_wallis_and_dunn(df):
    rows = []
    dunn_results = {}

    for variable in ANALYSIS_VARIABLES:
        groups = [
            df.loc[df["Upload source"] == source, variable].dropna().to_numpy()
            for source in SOURCE_ORDER
        ]
        statistic, p_value = kruskal(*groups)
        rows.append(
            {
                "Variable": variable,
                "Kruskal-Wallis H": statistic,
                "Degrees of freedom": len(groups) - 1,
                "P value": p_value,
            }
        )

        if p_value < 0.05:
            dunn_results[variable] = sp.posthoc_dunn(
                df,
                val_col=variable,
                group_col="Upload source",
                p_adjust="holm",
                sort=False,
            ).reindex(index=SOURCE_ORDER, columns=SOURCE_ORDER)

    return pd.DataFrame(rows), dunn_results


def spearman_with_holm(df):
    variables = ANALYSIS_VARIABLES
    rho_matrix = pd.DataFrame(
        np.eye(len(variables)),
        index=variables,
        columns=variables,
        dtype=float,
    )
    raw_p_matrix = pd.DataFrame(
        np.zeros((len(variables), len(variables))),
        index=variables,
        columns=variables,
        dtype=float,
    )
    adjusted_p_matrix = raw_p_matrix.copy()

    pairs = []
    p_values = []

    for i, first in enumerate(variables):
        for j in range(i + 1, len(variables)):
            second = variables[j]
            rho, p_value = spearmanr(
                df[first].to_numpy(),
                df[second].to_numpy(),
            )
            rho_matrix.loc[first, second] = rho
            rho_matrix.loc[second, first] = rho
            raw_p_matrix.loc[first, second] = p_value
            raw_p_matrix.loc[second, first] = p_value
            pairs.append((first, second))
            p_values.append(p_value)

    _, adjusted_p_values = pg.multicomp(
        p_values,
        method="holm",
    )

    for (first, second), adjusted_p in zip(pairs, adjusted_p_values):
        adjusted_p_matrix.loc[first, second] = adjusted_p
        adjusted_p_matrix.loc[second, first] = adjusted_p

    return rho_matrix, raw_p_matrix, adjusted_p_matrix


def inter_rater_agreement(df):
    rows = []

    measures = {
        "DISCERN": ("DISCERN rater 1", "DISCERN rater 2"),
        "JAMA": ("JAMA rater 1", "JAMA rater 2"),
        "GQS": ("GQS rater 1", "GQS rater 2"),
    }

    for measure, (rater_1, rater_2) in measures.items():
        long_data = pd.DataFrame(
            {
                "Video ID": np.repeat(df["Video ID"].to_numpy(), 2),
                "Rater": np.tile(["Rater 1", "Rater 2"], len(df)),
                "Score": np.column_stack(
                    [df[rater_1].to_numpy(), df[rater_2].to_numpy()]
                ).reshape(-1),
            }
        )

        icc_table = pg.intraclass_corr(
            data=long_data,
            targets="Video ID",
            raters="Rater",
            ratings="Score",
        )
        icc2 = icc_table.loc[icc_table["Type"] == "ICC2"].iloc[0]

        kappa = cohen_kappa_score(
            df[rater_1],
            df[rater_2],
            weights="quadratic",
        )

        rows.append(
            {
                "Assessment tool": measure,
                "ICC(2,1)": icc2["ICC"],
                "ICC 95% CI": str(icc2["CI95%"]),
                "Quadratic weighted kappa": kappa,
            }
        )

    return pd.DataFrame(rows)


def top_ten_videos(df):
    columns = [
        "Video ID",
        "Video title",
        "Video URL",
        "Upload source",
        "Views",
        "Likes",
        "Comments",
        "Country",
        "Mean DISCERN",
        "Mean JAMA",
        "Mean GQS",
        "Composite quality score",
    ]
    return (
        df.sort_values(
            ["Composite quality score", "Video ID"],
            ascending=[False, True],
        )
        .loc[:, columns]
        .head(10)
    )


def save_outputs(
    descriptive,
    by_source,
    kruskal_results,
    dunn_results,
    rho_matrix,
    raw_p_matrix,
    adjusted_p_matrix,
    agreement,
    top_ten,
):
    OUTPUT_DIR.mkdir(exist_ok=True)

    descriptive.to_csv(
        OUTPUT_DIR / "descriptive_statistics.csv",
        float_format="%.6f",
    )
    by_source.to_csv(
        OUTPUT_DIR / "summary_by_upload_source.csv",
        float_format="%.6f",
    )
    kruskal_results.to_csv(
        OUTPUT_DIR / "kruskal_wallis_results.csv",
        index=False,
        float_format="%.10g",
    )

    for variable, matrix in dunn_results.items():
        safe_name = (
            variable.lower()
            .replace(" ", "_")
            .replace("(", "")
            .replace(")", "")
        )
        matrix.to_csv(
            OUTPUT_DIR / f"dunn_holm_{safe_name}.csv",
            float_format="%.10g",
        )

    rho_matrix.to_csv(
        OUTPUT_DIR / "spearman_rho_matrix.csv",
        float_format="%.6f",
    )
    raw_p_matrix.to_csv(
        OUTPUT_DIR / "spearman_raw_p_matrix.csv",
        float_format="%.10g",
    )
    adjusted_p_matrix.to_csv(
        OUTPUT_DIR / "spearman_holm_adjusted_p_matrix.csv",
        float_format="%.10g",
    )
    agreement.to_csv(
        OUTPUT_DIR / "inter_rater_agreement.csv",
        index=False,
        float_format="%.6f",
    )
    top_ten.to_csv(
        OUTPUT_DIR / "top_10_composite_quality.csv",
        index=False,
        float_format="%.2f",
    )


def main():
    df = load_data()

    descriptive = descriptive_statistics(df)
    by_source = source_summary(df)
    kruskal_results, dunn_results = kruskal_wallis_and_dunn(df)
    rho_matrix, raw_p_matrix, adjusted_p_matrix = spearman_with_holm(df)
    agreement = inter_rater_agreement(df)
    top_ten = top_ten_videos(df)

    save_outputs(
        descriptive,
        by_source,
        kruskal_results,
        dunn_results,
        rho_matrix,
        raw_p_matrix,
        adjusted_p_matrix,
        agreement,
        top_ten,
    )

    print(f"Python: {platform.python_version()}")
    print(f"NumPy: {np.__version__}")
    print(f"pandas: {pd.__version__}")
    print(f"SciPy: {scipy.__version__}")
    print(f"scikit-posthocs: {sp.__version__}")
    print(f"Pingouin: {pg.__version__}")
    print(f"scikit-learn: {sklearn.__version__}")
    print()

    for variable in ["Mean DISCERN", "Mean JAMA", "Mean GQS"]:
        print(
            f"{variable}: "
            f"{round_half_up(descriptive.loc[variable, 'Mean'])} ± "
            f"{round_half_up(descriptive.loc[variable, 'SD'])}"
        )

    print()
    print(kruskal_results.to_string(index=False))
    print()
    print(agreement.to_string(index=False))
    print()
    print(top_ten.to_string(index=False))


if __name__ == "__main__":
    main()
