#!/usr/bin/env python3
"""
Supplementary File 2. Python analytic code for:
Predictors of Patient Satisfaction in Three Rural Districts of Sierra Leone:
A Cross-sectional Household Survey

Purpose:
    Reproduce the main descriptive statistics, Table 1, and multivariable ordinal
    logistic regression used for the patient satisfaction analysis.

Input:
    household_survey_final.xlsx
    Sheet name: Database

Outputs:
    outputs/table1_participant_characteristics.csv
    outputs/table2_ordinal_logistic_regression.csv
    outputs/figure1_forest_plot.png
    outputs/analysis_dataset.csv

Software:
    Python 3.11+
    pandas, numpy, scipy, statsmodels, matplotlib
"""

from pathlib import Path
import warnings
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from scipy import stats
from statsmodels.miscmodels.ordinal_model import OrderedModel

warnings.filterwarnings("ignore", category=FutureWarning)

DATA_PATH = Path("household_survey_final.xlsx")
SHEET_NAME = "Database"
OUTPUT_DIR = Path("outputs")
OUTPUT_DIR.mkdir(exist_ok=True)

REQUIRED_COLUMNS = [
    "district", "chiefdom", "respondent_age", "respondent_gender",
    "marital_status", "education_level", "visited_health_facility",
    "visit_frequency", "facility_most_visited", "reason_for_visit",
    "transport_mode", "travel_time", "overall_satisfaction",
    "staff_attitude_satisfaction", "medicine_availability_satisfaction",
    "waiting_time_satisfaction"
]

def clean_string_columns(data):
    """Standardise string variables to lower-case trimmed values."""
    out = data.copy()
    for col in out.select_dtypes(include="object").columns:
        out[col] = out[col].astype(str).str.strip().str.lower()
        out.loc[out[col].isin(["nan", "none", ""]), col] = np.nan
    return out

def validate_columns(data, required_columns):
    missing = [col for col in required_columns if col not in data.columns]
    if missing:
        raise ValueError(f"Missing required columns: {missing}")

def prepare_analysis_dataset(data):
    """Prepare respondent-level analytic dataset."""
    data = clean_string_columns(data)
    validate_columns(data, REQUIRED_COLUMNS)

    data["analytic_sample"] = (
        (data["visited_health_facility"] == "yes") &
        data["overall_satisfaction"].notna()
    ).astype(int)

    ana = data.loc[data["analytic_sample"] == 1].copy()

    ana["female"] = (ana["respondent_gender"] == "female").astype(int)
    ana["satisfied_or_very_satisfied"] = (ana["overall_satisfaction"] >= 4).astype(int)

    travel_order = {
        "less_than_30": 1,
        "30_to_1_hour": 2,
        "1_to_2_hours": 3,
        "more_than_2_hours": 4
    }
    ana["travel_time_category"] = ana["travel_time"].map(travel_order)

    ana["hospital_or_other_non_chc"] = (
        ana["facility_most_visited"] != "community_health_center"
    ).astype(int)

    ana["district_kailahun"] = (ana["district"] == "kailahun").astype(int)
    ana["district_pujehun"] = (ana["district"] == "pujehun").astype(int)

    # Age bands used for descriptive Table 1.
    ana["age_group"] = pd.cut(
        ana["respondent_age"],
        bins=[13, 17, 24, 34, 49, np.inf],
        labels=["14-17 years", "18-24 years", "25-34 years", "35-49 years", ">=50 years"],
        right=True
    )

    return ana

def n_pct(series):
    counts = series.value_counts(dropna=False)
    total = len(series)
    return pd.DataFrame({
        "category": counts.index.astype(str),
        "n": counts.values,
        "percent": np.round(counts.values / total * 100, 1)
    })

def table1_by_district(ana):
    """Generate core participant characteristics by district."""
    rows = []

    def add_continuous(var, label):
        overall = ana[var].dropna()
        rows.append({
            "characteristic": f"{label}, mean (SD)",
            "overall": f"{overall.mean():.1f} ({overall.std():.1f})",
            "kailahun": fmt_mean_sd(ana.loc[ana.district == "kailahun", var]),
            "kambia": fmt_mean_sd(ana.loc[ana.district == "kambia", var]),
            "pujehun": fmt_mean_sd(ana.loc[ana.district == "pujehun", var]),
        })
        rows.append({
            "characteristic": f"{label}, median (IQR)",
            "overall": fmt_median_iqr(overall),
            "kailahun": fmt_median_iqr(ana.loc[ana.district == "kailahun", var]),
            "kambia": fmt_median_iqr(ana.loc[ana.district == "kambia", var]),
            "pujehun": fmt_median_iqr(ana.loc[ana.district == "pujehun", var]),
        })

    def add_categorical(var, label, order=None):
        rows.append({"characteristic": label, "overall": "", "kailahun": "", "kambia": "", "pujehun": ""})
        categories = order if order is not None else sorted(ana[var].dropna().astype(str).unique())
        for cat in categories:
            rows.append({
                "characteristic": f"  {cat}",
                "overall": fmt_n_pct(ana[var], cat),
                "kailahun": fmt_n_pct(ana.loc[ana.district == "kailahun", var], cat),
                "kambia": fmt_n_pct(ana.loc[ana.district == "kambia", var], cat),
                "pujehun": fmt_n_pct(ana.loc[ana.district == "pujehun", var], cat),
            })

    add_continuous("respondent_age", "Age")
    add_categorical("age_group", "Age group", ["14-17 years", "18-24 years", "25-34 years", "35-49 years", ">=50 years"])
    add_categorical("respondent_gender", "Sex", ["female", "male"])
    add_categorical("education_level", "Education level", ["no_education", "primary", "secondary", "tertiary", "other"])
    add_categorical("marital_status", "Marital status", ["single", "married", "separated", "divorced", "widowed"])
    add_categorical("facility_most_visited", "Facility most frequently used")
    add_categorical("reason_for_visit", "Main reason for visit")
    add_categorical("transport_mode", "Main transport mode")
    add_categorical("travel_time", "Travel time to facility", ["less_than_30", "30_to_1_hour", "1_to_2_hours", "more_than_2_hours"])
    add_categorical("overall_satisfaction", "Single-item overall satisfaction", [1, 2, 3, 4, 5])
    add_categorical("satisfied_or_very_satisfied", "Satisfied or very satisfied", [1])

    table = pd.DataFrame(rows)
    table.to_csv(OUTPUT_DIR / "table1_participant_characteristics.csv", index=False)
    return table

def fmt_mean_sd(x):
    x = pd.to_numeric(x, errors="coerce").dropna()
    return f"{x.mean():.1f} ({x.std():.1f})"

def fmt_median_iqr(x):
    x = pd.to_numeric(x, errors="coerce").dropna()
    return f"{x.median():.1f} ({x.quantile(0.25):.1f}-{x.quantile(0.75):.1f})"

def fmt_n_pct(series, category):
    s = series.astype(str)
    c = str(category)
    n = (s == c).sum()
    denom = s.notna().sum()
    return f"{n} ({n / denom * 100:.1f})" if denom > 0 else "0 (0.0)"

def fit_ordinal_model(ana):
    """
    Fit multivariable ordinal logistic regression.
    Outcome: five-level overall satisfaction.
    Reference district: Kambia.
    """
    model_vars = [
        "overall_satisfaction",
        "staff_attitude_satisfaction",
        "waiting_time_satisfaction",
        "medicine_availability_satisfaction",
        "female",
        "visit_frequency",
        "travel_time_category",
        "respondent_age",
        "hospital_or_other_non_chc",
        "district_kailahun",
        "district_pujehun",
        "chiefdom"
    ]
    m = ana[model_vars].dropna().copy()
    y = m["overall_satisfaction"].astype(int)

    X = m[
        [
            "staff_attitude_satisfaction",
            "waiting_time_satisfaction",
            "medicine_availability_satisfaction",
            "female",
            "visit_frequency",
            "travel_time_category",
            "respondent_age",
            "hospital_or_other_non_chc",
            "district_kailahun",
            "district_pujehun"
        ]
    ]

    ordered_model = OrderedModel(y, X, distr="logit")
    result = ordered_model.fit(method="bfgs", disp=False, maxiter=1000)

    # Cluster-robust standard errors at chiefdom level.
    # statsmodels stores the updated covariance in the same result object.
    result._get_robustcov_results(cov_type="cluster", groups=m["chiefdom"])

    coef_names = X.columns.tolist()
    estimates = []
    for name in coef_names:
        beta = result.params[name]
        se = result.bse[name]
        z = beta / se
        p = 2 * (1 - stats.norm.cdf(abs(z)))
        estimates.append({
            "predictor": name,
            "beta": beta,
            "standard_error": se,
            "adjusted_odds_ratio": np.exp(beta),
            "lower_95_ci": np.exp(beta - 1.96 * se),
            "upper_95_ci": np.exp(beta + 1.96 * se),
            "p_value": p
        })

    table2 = pd.DataFrame(estimates)
    table2.to_csv(OUTPUT_DIR / "table2_ordinal_logistic_regression.csv", index=False)
    return result, table2

def make_forest_plot(table2):
    """Create forest plot from model-derived estimates."""
    label_map = {
        "staff_attitude_satisfaction": "Staff attitude",
        "waiting_time_satisfaction": "Waiting-time satisfaction",
        "medicine_availability_satisfaction": "Medicine availability",
        "female": "Female sex",
        "visit_frequency": "Visit frequency",
        "travel_time_category": "Travel-time category",
        "respondent_age": "Age",
        "hospital_or_other_non_chc": "Hospital or other non-CHC facility",
        "district_kailahun": "Kailahun vs Kambia",
        "district_pujehun": "Pujehun vs Kambia"
    }

    d = table2.copy()
    d["label"] = d["predictor"].map(label_map)
    d = d.iloc[::-1].reset_index(drop=True)

    y = np.arange(len(d))
    fig, ax = plt.subplots(figsize=(10, 6), dpi=300)
    ax.errorbar(
        d["adjusted_odds_ratio"], y,
        xerr=[
            d["adjusted_odds_ratio"] - d["lower_95_ci"],
            d["upper_95_ci"] - d["adjusted_odds_ratio"]
        ],
        fmt="o", capsize=4, linewidth=1.2
    )
    ax.axvline(1, linestyle="--", linewidth=1)
    ax.set_xscale("log")
    ax.set_yticks(y)
    ax.set_yticklabels(d["label"])
    ax.set_xlabel("Adjusted odds ratio, log scale")
    ax.set_title("Predictors of Higher Patient Satisfaction")
    fig.tight_layout()
    fig.savefig(OUTPUT_DIR / "figure1_forest_plot.png", dpi=300)
    fig.savefig(OUTPUT_DIR / "figure1_forest_plot.pdf")
    plt.close(fig)

def main():
    raw = pd.read_excel(DATA_PATH, sheet_name=SHEET_NAME)
    ana = prepare_analysis_dataset(raw)
    ana.to_csv(OUTPUT_DIR / "analysis_dataset.csv", index=False)

    print(f"Total surveyed respondents: {len(raw)}")
    print(f"Analytic sample for satisfaction analysis: {len(ana)}")

    table1_by_district(ana)
    result, table2 = fit_ordinal_model(ana)
    make_forest_plot(table2)

    print("Analysis complete. Outputs saved to:", OUTPUT_DIR.resolve())

if __name__ == "__main__":
    main()
