import os
import re
import numpy as np
import pandas as pd
from scipy.stats import chi2_contingency, fisher_exact, chi2
import matplotlib.pyplot as plt
import matplotlib as mpl
import seaborn as sns
import statsmodels.api as sm
from lifelines import CoxPHFitter
from PIL import Image, ImageDraw, ImageFont

INPUT_PATH = r"G:\Influenza_Reinfection\VSCODE\Influenza_Merged_Data.csv"
OUT_DIR = r"G:\Influenza_Reinfection\Output_Data"
KDE_INPUT_PATH = r"G:\Influenza_Reinfection\Output_Data\Influenza_Repeated_Patients_Excluded_14d.csv"

os.makedirs(OUT_DIR, exist_ok=True)

print("INPUT:", INPUT_PATH)
print("OUT  :", OUT_DIR)

NEEDED = {
    "ID_number": ["ID_number", "ID", "Identification", "Valid_ID"],
    "Sex": ["Sex", "Gender"],
    "Age": ["Age", "Age_years"],
    "Address": ["Address", "Current_address", "Detailed_address"],
    "Population_category": ["Population_category", "Occupation_category"],
    "Case_category": ["Case_category", "Diagnosis_type"],
    "Onset_date": ["Onset_date", "Date_of_onset", "Illness_date"],
    "Disease_name": ["Disease_name", "Disease", "Diagnosis"]
}

def clean_colname(c):
    return str(c).strip()

def pick_col(cols, candidates):
    for c in candidates:
        if c in cols:
            return c
    return None

def clean_id(x):
    if pd.isna(x):
        return pd.NA
    s = str(x).strip()
    if s in {".", "nan", "None", ""}:
        return pd.NA
    s = s.lstrip("'").strip()
    return s if s else pd.NA

def parse_date(s):
    if pd.isna(s):
        return pd.NaT
    ss = str(s).strip()
    if ss in {".", "nan", "None", ""}:
        return pd.NaT
    return pd.to_datetime(ss, errors="coerce")

def classify_city(addr):
    if pd.isna(addr):
        return "Urban"
    return "Rural" if "Town" in str(addr) else "Urban"

def classify_occupation(pop):
    if pd.isna(pop):
        return "Others"
    s = str(pop).strip()
    if s == "Student":
        return "Students"
    if s == "Preschool":
        return "Preschool-aged children"
    if s == "Medical":
        return "Medical staff"
    if s == "Farmer":
        return "Farmers"
    if s == "Teacher":
        return "Teachers"
    return "Others"

def classify_case(case_cls):
    if pd.isna(case_cls):
        return "Clinical diagnosis cases"
    return "laboratory-confirmed case" if str(case_cls).strip() == "Confirmed" else "Clinical diagnosis cases"

def age_to_years(age_str):
    if pd.isna(age_str):
        return pd.NA
    m = re.search(r"(\d+)", str(age_str))
    return int(m.group(1)) if m else pd.NA

def age_group(a):
    if pd.isna(a):
        return pd.NA
    a = int(a)
    if 0 <= a <= 5:   return "0-5"
    if 6 <= a <= 19:  return "6-19"
    if 20 <= a <= 45: return "20-45"
    if 46 <= a <= 60: return "46-60"
    if a >= 61:       return ">=61"
    return pd.NA

try:
    df_head = pd.read_csv(INPUT_PATH, nrows=5, dtype=str, engine="python")
    cols = [clean_colname(c) for c in df_head.columns.tolist()]

    colmap = {}
    for std, cand_list in NEEDED.items():
        found = pick_col(cols, [clean_colname(x) for x in cand_list])
        colmap[std] = found

    missing = [k for k,v in colmap.items() if v is None]
    if missing:
        raise KeyError("Missing critical columns. Please check input CSV headers.")

    usecols_real = [colmap[k] for k in colmap]

    chunks = []
    chunksize = 200000

    for i, ck in enumerate(pd.read_csv(
        INPUT_PATH,
        usecols=usecols_real,
        dtype=str,
        engine="python",
        chunksize=chunksize
    )):
        ck.columns = [clean_colname(c) for c in ck.columns]
        rename_dict = {colmap[k]: k for k in colmap}
        ck = ck.rename(columns=rename_dict)
        ck["Disease_name"] = ck["Disease_name"].astype(str).str.strip()
        ck = ck[ck["Disease_name"] == "Influenza"].copy()
        chunks.append(ck)

    df = pd.concat(chunks, ignore_index=True)

    df["ID_clean"] = df["ID_number"].map(clean_id)
    df["Onset_date_dt"] = df["Onset_date"].map(parse_date)
    df = df.dropna(subset=["ID_clean", "Onset_date_dt"]).copy()

    df["Sex_clean"] = df["Sex"].astype(str).str.strip()
    df.loc[~df["Sex_clean"].isin(["Male","Female"]), "Sex_clean"] = pd.NA

    df["Age_years"] = df["Age"].map(age_to_years)
    df["Age_group"] = df["Age_years"].map(age_group)
    df["City"] = df["Address"].map(classify_city)
    df["Occupation"] = df["Population_category"].map(classify_occupation)
    df["Case_classification"] = df["Case_category"].map(classify_case)

    df = df.sort_values(["ID_clean", "Onset_date_dt"]).copy()
    df["prev_date"] = df.groupby("ID_clean")["Onset_date_dt"].shift(1)
    df["gap_days"] = (df["Onset_date_dt"] - df["prev_date"]).dt.days
    df["is_new_episode"] = df["gap_days"].isna() | (df["gap_days"] > 14)
    df["episode_no"] = df.groupby("ID_clean")["is_new_episode"].cumsum()
    df["excluded_14d"] = ~df["is_new_episode"]

    episode_counts = (
        df[df["is_new_episode"]]
        .groupby("ID_clean")["episode_no"]
        .max()
        .rename("episode_count")
    )
    df = df.merge(episode_counts, on="ID_clean", how="left")

    table2 = df.copy()
    episode_start = (
        table2[table2["is_new_episode"]]
        .groupby(["ID_clean","episode_no"])["Onset_date_dt"]
        .min()
        .rename("episode_start_date")
        .reset_index()
    )
    table2 = table2.merge(episode_start, on=["ID_clean","episode_no"], how="left")

    OUT_TABLE2 = os.path.join(OUT_DIR, "Table_2_All_Cases_Episodes.csv")
    table2.to_csv(OUT_TABLE2, index=False)

    table3 = (
        df[(df["episode_no"] == 1) & (df["is_new_episode"])]
        .sort_values(["ID_clean","Onset_date_dt"])
        .copy()
    )
    OUT_TABLE3 = os.path.join(OUT_DIR, "Table_3_First_Infection.csv")
    table3.to_csv(OUT_TABLE3, index=False)

    episode_cols = ["1","2","3",">3"]

    def epi_bucket(k):
        if pd.isna(k): return pd.NA
        k = int(k)
        if k == 1: return "1"
        if k == 2: return "2"
        if k == 3: return "3"
        return ">3"

    def format_n_pct(n, denom):
        if denom == 0:
            return f"{n} (0.0)"
        return f"{n} ({n/denom*100:.1f})"

    def chisq_or_fisher(table):
        arr = np.array(table, dtype=float)
        if (arr.sum(axis=1) == 0).any() or (arr.sum(axis=0) == 0).any():
            return np.nan, np.nan, "NA"
        if arr.shape == (2,2):
            chi2_val, p, dof, exp = chi2_contingency(arr, correction=False)
            if (exp < 5).any():
                _, pf = fisher_exact(arr.astype(int))
                return np.nan, pf, "Fisher"
            return chi2_val, p, "Chi-square"
        chi2_val, p, dof, exp = chi2_contingency(arr, correction=False)
        return chi2_val, p, "Chi-square"

    def median_iqr(series):
        s = pd.to_numeric(series, errors="coerce").dropna()
        if len(s)==0:
            return ""
        q1 = s.quantile(0.25)
        med= s.quantile(0.50)
        q3 = s.quantile(0.75)
        return f"{med:.1f} ({q1:.1f}-{q3:.1f})"

    first_episode = table3.copy()
    first_episode["episode_bucket"] = first_episode["episode_count"].map(epi_bucket)

    N_total = first_episode["ID_clean"].nunique()
    col_denoms = first_episode["episode_bucket"].value_counts().reindex(episode_cols).fillna(0).astype(int).to_dict()

    table1_rows = []
    table1_rows.append({
        "Characteristic": "Total",
        "Total": f"{N_total}",
        "Patients with 1 episode, N (%)": format_n_pct(col_denoms["1"], N_total),
        "Patients with 2 episode, N (%)": format_n_pct(col_denoms["2"], N_total),
        "Patients with 3 episode, N (%)": format_n_pct(col_denoms["3"], N_total),
        "Patients with >3 episodes(reinfection), N (%)": format_n_pct(col_denoms[">3"], N_total),
        "Chi-square": "",
        "P-value": ""
    })

    def build_block(df_person, var, levels, label):
        rows = []
        den = df_person["episode_bucket"].value_counts().reindex(episode_cols).fillna(0).astype(int).to_dict()
        total = df_person.shape[0]

        ct = pd.crosstab(df_person[var], df_person["episode_bucket"]).reindex(index=levels, columns=episode_cols).fillna(0).astype(int)
        chi2_val, p, method = chisq_or_fisher(ct.values)

        rows.append({
            "Characteristic": label,
            "Total": "",
            "Patients with 1 episode, N (%)": "",
            "Patients with 2 episode, N (%)": "",
            "Patients with 3 episode, N (%)": "",
            "Patients with >3 episodes(reinfection), N (%)": "",
            "Chi-square": (f"{chi2_val:.3f}" if method=="Chi-square" and pd.notna(chi2_val) else ""),
            "P-value": (f"{p:.4f}" if pd.notna(p) else "")
        })

        for lv in levels:
            sub = df_person[df_person[var] == lv]
            nT = len(sub)
            n1 = int((sub["episode_bucket"]=="1").sum())
            n2 = int((sub["episode_bucket"]=="2").sum())
            n3 = int((sub["episode_bucket"]=="3").sum())
            n4 = int((sub["episode_bucket"]==">3").sum())

            rows.append({
                "Characteristic": f"  {lv}",
                "Total": format_n_pct(nT, total),
                "Patients with 1 episode, N (%)": format_n_pct(n1, den["1"]),
                "Patients with 2 episode, N (%)": format_n_pct(n2, den["2"]),
                "Patients with 3 episode, N (%)": format_n_pct(n3, den["3"]),
                "Patients with >3 episodes(reinfection), N (%)": format_n_pct(n4, den[">3"]),
                "Chi-square": "",
                "P-value": ""
            })
        return rows

    sex_df = first_episode.dropna(subset=["Sex_clean"]).copy()
    table1_rows += build_block(sex_df, "Sex_clean", ["Male","Female"], "Sex")

    age_df = first_episode.dropna(subset=["Age_group"]).copy()
    table1_rows += build_block(age_df, "Age_group", ["0-5","6-19","20-45","46-60",">=61"], "Age at first episode (years)")

    table1_rows.append({
        "Characteristic": "Median (IQR*)",
        "Total": median_iqr(first_episode["Age_years"]),
        "Patients with 1 episode, N (%)": median_iqr(first_episode.loc[first_episode["episode_bucket"]=="1","Age_years"]),
        "Patients with 2 episode, N (%)": median_iqr(first_episode.loc[first_episode["episode_bucket"]=="2","Age_years"]),
        "Patients with 3 episode, N (%)": median_iqr(first_episode.loc[first_episode["episode_bucket"]=="3","Age_years"]),
        "Patients with >3 episodes(reinfection), N (%)": median_iqr(first_episode.loc[first_episode["episode_bucket"]==">3","Age_years"]),
        "Chi-square": "",
        "P-value": ""
    })

    table1_rows += build_block(first_episode, "Occupation",
                               ["Farmers","Students","Teachers","Preschool-aged children","Medical staff","Others"],
                               "Occupation")

    table1_rows += build_block(first_episode, "Case_classification",
                               ["Clinical diagnosis cases","laboratory-confirmed case"],
                               "Case classification")

    table1_rows += build_block(first_episode, "City", ["Urban","Rural"], "City")

    table1 = pd.DataFrame(table1_rows)
    OUT_TABLE1 = os.path.join(OUT_DIR, "Table_1_Reinfection_Characteristics.csv")
    table1.to_csv(OUT_TABLE1, index=False)

    def to_bool(x):
        if pd.isna(x): 
            return np.nan
        s = str(x).strip().lower()
        if s in ["true","1","t","yes"]:
            return True
        if s in ["false","0","f","no"]:
            return False
        return np.nan

    t2 = table2.copy()
    t2["is_new_episode"] = t2["is_new_episode"].map(to_bool)
    t2["episode_no"] = pd.to_numeric(t2["episode_no"], errors="coerce")
    t2["episode_count"] = pd.to_numeric(t2["episode_count"], errors="coerce")
    t2["episode_start_date"] = pd.to_datetime(t2["episode_start_date"], errors="coerce")
    t2["ID_clean"] = t2["ID_clean"].astype(str).str.strip()

    all_people = (
        t2.dropna(subset=["ID_clean"])
          .drop_duplicates("ID_clean")[["ID_clean"]]
          .copy()
    )

    epi_starts = (
        t2[t2["is_new_episode"] == True]
        .dropna(subset=["ID_clean", "episode_start_date"])
        .sort_values(["ID_clean", "episode_no", "episode_start_date"])
        .loc[:, ["ID_clean","episode_no","episode_start_date",
                 "Sex_clean","Age_years","Age_group","Occupation","City","Case_classification","episode_count"]]
        .rename(columns={"episode_start_date":"episode_start"})
        .copy()
    )

    baseline = (
        epi_starts.sort_values(["ID_clean","episode_start"])
        .drop_duplicates("ID_clean", keep="first")
        .set_index("ID_clean")[["Sex_clean","Age_years","Age_group","Occupation","City","Case_classification","episode_count"]]
        .copy()
    )

    t0 = epi_starts.groupby("ID_clean")["episode_start"].min().rename("t0")
    event_date = epi_starts[epi_starts["episode_no"] == 2].groupby("ID_clean")["episode_start"].min().rename("event_date")
    censor_date = epi_starts.groupby("ID_clean")["episode_start"].max().rename("censor_date")
    study_end = epi_starts["episode_start"].max()

    surv = all_people.set_index("ID_clean").copy()
    surv = surv.join(baseline, how="left")
    surv = surv.join(t0, how="left")
    surv = surv.join(event_date, how="left")
    surv = surv.join(censor_date, how="left")

    surv["event"] = surv["event_date"].notna().astype(int)
    surv["t1"] = surv["event_date"].fillna(surv["censor_date"])
    surv["followup_days"] = (surv["t1"] - surv["t0"]).dt.days
    surv = surv.dropna(subset=["t0","t1","followup_days"]).copy()
    surv.loc[surv["followup_days"] <= 0, "followup_days"] = 1
    surv["followup_weeks"] = (surv["followup_days"] / 7.0).astype("float32")
    surv["person_years"] = surv["followup_days"] / 365.25

    OKABE_ITO = ["#000000", "#E69F00", "#56B4E9", "#009E73", "#F0E442", "#0072B2", "#D55E00", "#CC79A7"]

    def set_academic_style():
        mpl.rcParams.update({
            "figure.dpi": 120,
            "savefig.dpi": 300,
            "font.size": 12,
            "axes.titlesize": 16,
            "axes.labelsize": 13,
            "legend.fontsize": 11,
            "axes.linewidth": 1.0,
            "xtick.major.width": 1.0,
            "ytick.major.width": 1.0,
            "grid.alpha": 0.3,
            "font.family": "Arial",
            "axes.unicode_minus": False
        })
    set_academic_style()

    MAX_WEEKS = 156
    TIME_POINTS = [0, 26, 52, 78, 104, 130, 156]
    LINE_W = 4.25

    def format_risk_table_line(label, values, label_width=22, col_width=10):
        lab = str(label)
        return f"{lab:<{label_width}}" + "".join(f"{v:>{col_width}d}" for v in values)

    def wrap_label(s, width=26):
        s = str(s)
        if len(s) <= width:
            return s
        parts = s.replace("-", " - ").split()
        lines, cur = [], ""
        for p in parts:
            add = (cur + " " + p).strip()
            if len(add) <= width:
                cur = add
            else:
                if cur: lines.append(cur)
                cur = p
        if cur: lines.append(cur)
        return "\n".join(lines)

    def km_cuminc_curve(dur, evt, max_weeks=156):
        dur = np.minimum(dur.astype("float32", copy=False), max_weeks)
        evt = evt.astype("int8", copy=False)
        event_times = np.sort(np.unique(dur[(evt == 1) & (dur > 0) & (dur <= max_weeks)]))
        if event_times.size == 0:
            return np.array([0, max_weeks], dtype="float32"), np.array([0.0, 0.0], dtype="float32")
        n_at_t = np.array([np.sum(dur >= t) for t in event_times], dtype="int32")
        d_at_t = np.array([np.sum((dur == t) & (evt == 1)) for t in event_times], dtype="int32")
        hazard = d_at_t / np.maximum(n_at_t, 1)
        S = np.cumprod(1.0 - hazard)
        CI = 1.0 - S
        t_curve = np.concatenate(([0.0], event_times.astype("float32"), [max_weeks])).astype("float32")
        ci_curve = np.concatenate(([0.0], CI.astype("float32"), [CI[-1]])).astype("float32")
        return t_curve, ci_curve

    def cuminc_at_times(t_curve, ci_curve, times):
        times = np.asarray(times, dtype="float32")
        out = []
        for t in times:
            idx = np.searchsorted(t_curve, t, side="right") - 1
            if idx < 0: out.append(0.0)
            else: out.append(float(ci_curve[idx]))
        return out

    def logrank_multigroup(dur, evt, grp, max_weeks=156):
        dur = np.minimum(dur.astype("float32", copy=False), max_weeks)
        evt = evt.astype("int8", copy=False)
        mask = (~np.isnan(dur)) & (dur >= 0) & (dur <= max_weeks)
        dur = dur[mask]; evt = evt[mask]; grp = grp[mask]
        event_times = np.sort(np.unique(dur[evt == 1]))
        groups = pd.unique(grp)
        groups = [g for g in groups if pd.notna(g)]
        k = len(groups)
        if k <= 1 or event_times.size == 0:
            return np.nan, k-1, np.nan
        idx_by_g = {g: np.where(grp == g)[0] for g in groups}
        O = np.zeros(k, dtype="float64")
        E = np.zeros(k, dtype="float64")
        V = np.zeros((k, k), dtype="float64")
        for t in event_times:
            risk = (dur >= t)
            n = int(np.sum(risk))
            if n <= 1: continue
            event = (dur == t) & (evt == 1)
            d = int(np.sum(event))
            if d == 0: continue
            n_g = np.zeros(k, dtype="int32")
            d_g = np.zeros(k, dtype="int32")
            for i, g in enumerate(groups):
                gi = idx_by_g[g]
                n_g[i] = int(np.sum(risk[gi]))
                d_g[i] = int(np.sum(event[gi]))
            e_g = d * (n_g / n)
            O += d_g
            E += e_g
            v_factor = d * (n - d) / (n * (n - 1))
            for i in range(k):
                for j in range(k):
                    if i == j:
                        V[i, j] += v_factor * (n_g[i] - (n_g[i] * n_g[j] / n))
                    else:
                        V[i, j] += v_factor * (0 - (n_g[i] * n_g[j] / n))
        O1, E1, V1 = O[:-1], E[:-1], V[:-1, :-1]
        try:
            stat = float((O1 - E1).T @ np.linalg.pinv(V1) @ (O1 - E1))
            df = k - 1
            p = float(1.0 - chi2.cdf(stat, df))
            return stat, df, p
        except Exception:
            return np.nan, k-1, np.nan

    def normalize_age_label(x: str) -> str:
        s = str(x).strip()
        s = s.replace(">=", ">=")
        return s

    def age_sort_key(x: str) -> int:
        s = normalize_age_label(x)
        if s in ["0-5"]: return 0
        if s in ["6-19"]: return 1
        if s in ["20-45"]: return 2
        if s in ["46-60"]: return 3
        if s in [">=61", "61+", "60+"]: return 4
        m = re.search(r"\d+", s)
        return int(m.group()) if m else 999

    def normalize_gender_label(x: str) -> str:
        s = str(x).strip()
        if s in ["Male", "M", "male"]: return "Male"
        if s in ["Female", "F", "female"]: return "Female"
        if s == "" or s.lower() in ["nan", "none"]: return "Unknown"
        return s

    def display_group_name(group_col: str) -> str:
        if group_col in ["Age_group"]: return "Age group"
        if group_col in ["Sex_clean"]: return "Gender group"
        return group_col

    def preprocess_group_series(group_col: str, s: pd.Series) -> pd.Series:
        if group_col in ["Age_group"]: return s.map(normalize_age_label)
        if group_col in ["Sex_clean"]: return s.map(normalize_gender_label)
        return s.astype(str)

    def get_levels_sorted(group_col: str, vc: pd.Series) -> list:
        levels = vc.index.tolist()
        if group_col in ["Age_group"]:
            levels = sorted(levels, key=age_sort_key)
        if group_col in ["Sex_clean"]:
            pref = {"Male": 0, "Female": 1, "Unknown": 2}
            levels = sorted(levels, key=lambda x: pref.get(str(x), 99))
        return levels

    def plot_group_km_academic(surv_df, group_col, out_png, min_group_n=30, save_ci_csv=True):
        df = surv_df.copy()
        if group_col == "Occupation":
            exclude_list = ["Students", "Preschool-aged children"]
            df = df[~df[group_col].isin(exclude_list)]
            
        df[group_col] = preprocess_group_series(group_col, df[group_col])
        df = df.dropna(subset=[group_col]).copy()
        vc = df[group_col].value_counts()
        vc = vc[vc >= min_group_n]
        levels = get_levels_sorted(group_col, vc)

        if len(levels) < 2: return None

        dur_all = pd.to_numeric(df["followup_weeks"], errors="coerce").astype("float32").to_numpy()
        evt_all = pd.to_numeric(df["event"], errors="coerce").fillna(0).astype("int8").to_numpy()
        grp_all = df[group_col].astype(str).to_numpy()
        stat, dflr, p = logrank_multigroup(dur_all, evt_all, grp_all, MAX_WEEKS)

        curves = {}
        max_ci = 0.0
        ci_rows = []

        for lv in levels:
            sub = df[df[group_col] == lv]
            dur = pd.to_numeric(sub["followup_weeks"], errors="coerce").astype("float32").to_numpy()
            evt = pd.to_numeric(sub["event"], errors="coerce").fillna(0).astype("int8").to_numpy()
            t_curve, ci_curve = km_cuminc_curve(dur, evt, MAX_WEEKS)
            dur_clip = np.minimum(dur, MAX_WEEKS)
            curves[lv] = (t_curve, ci_curve, dur_clip)
            max_ci = max(max_ci, float(np.nanmax(ci_curve)))
            cis = cuminc_at_times(t_curve, ci_curve, TIME_POINTS)
            for w, v in zip(TIME_POINTS, cis):
                ci_rows.append({
                    "group_var": display_group_name(group_col),
                    "level": lv,
                    "week": int(w),
                    "cuminc_1_minus_S": float(v),
                    "cuminc_percent": float(v * 100.0)
                })

        y_top = min(1.0, max(0.02, max_ci * 1.08))
        title_name = display_group_name(group_col)
        ci_df_print = pd.DataFrame(ci_rows)

        if save_ci_csv:
            out_ci_csv = os.path.splitext(out_png)[0] + "_CumInc_keyweeks.csv"
            ci_df_print.to_csv(out_ci_csv, index=False)

        fig = plt.figure(figsize=(12.5, 8.5))
        gs = fig.add_gridspec(2, 1, height_ratios=[4, 1.45], hspace=0.20)
        ax = fig.add_subplot(gs[0, 0])
        ax_tab = fig.add_subplot(gs[1, 0])
        ax_tab.axis("off")

        for i, lv in enumerate(levels):
            t_curve, ci_curve, _ = curves[lv]
            color = OKABE_ITO[(i + 1) % len(OKABE_ITO)]
            ax.step(t_curve, ci_curve, where="post", linewidth=LINE_W, color=color, label=wrap_label(lv, 26))

        ax.set_xlim(0, MAX_WEEKS)
        ax.set_xticks(TIME_POINTS)
        ax.set_ylim(0, y_top)
        ax.set_xlabel("Weeks after primary episode")
        ax.set_ylabel("Probability of Reinfection")
        ax.set_title(f"{title_name} (log-rank p={p:.3g})")
        ax.grid(True)

        leg = ax.legend(title=title_name, frameon=True, ncol=2, loc="upper left")
        leg.get_frame().set_alpha(0.9)

        label_w = 22
        col_w = 10
        ax_tab.text(0.01, 0.82, "Number at risk", fontsize=14, fontweight="bold")
        ax_tab.text(0.01, 0.58, format_risk_table_line("Weeks:", TIME_POINTS, label_width=label_w, col_width=col_w), family="monospace")

        y = 0.28
        dy = 0.12 if len(levels) <= 6 else 0.08
        for lv in levels:
            _, _, dur_lv = curves[lv]
            ar = [int(np.sum(dur_lv >= t)) for t in TIME_POINTS]
            ax_tab.text(0.01, y, format_risk_table_line(str(lv), ar, label_width=label_w, col_width=col_w), family="monospace")
            y -= dy

        fig.savefig(out_png, dpi=300, bbox_inches="tight")
        plt.close(fig)
        return {"Group": title_name, "chi2": stat, "df": dflr, "p": p, "png": out_png}

    GROUP_VARS = ["Sex_clean", "Age_group", "Occupation", "City", "Case_classification"]
    results = []
    for g in GROUP_VARS:
        title_name = display_group_name(g).replace(" ", "_")
        out_png = os.path.join(OUT_DIR, f"KM_{title_name}_Weeks_0_156_Academic_EN.png")
        r = plot_group_km_academic(surv, g, out_png, min_group_n=30, save_ci_csv=True)
        if r:
            results.append(r)

    logrank_summary = pd.DataFrame(results)
    out_csv = os.path.join(OUT_DIR, "Logrank_summary_Weeks_0_156_Academic_EN.csv")
    logrank_summary.to_csv(out_csv, index=False)

    def combine_images_lowercase():
        INPUT_PATHS = [
            os.path.join(OUT_DIR, "KM_Overall"),
            os.path.join(OUT_DIR, "KM_Occupation_Weeks_0_156_Academic_EN"),
            os.path.join(OUT_DIR, "KM_Gender_group_Weeks_0_156_Academic_EN"),
            os.path.join(OUT_DIR, "KM_City_Weeks_0_156_Academic_EN"),
            os.path.join(OUT_DIR, "KM_Case_classification_Weeks_0_156_Academic_EN"),
            os.path.join(OUT_DIR, "KM_Age_group_Weeks_0_156_Academic_EN")
        ]
        OUTPUT_NAME = "Combined_KM_Panels_Lowercase"
        COLS, ROWS = 3, 2
        LABEL_SIZE, PADDING, LABEL_OFFSET = 130, 40, 30

        loaded_images = []
        max_w, max_h = 0, 0
        for path in INPUT_PATHS:
            valid_path = None
            for ext in [".png", ".PNG", ".tif", ".TIF"]:
                full_path = path + ext if not path.lower().endswith(ext.lower()) else path
                if os.path.exists(full_path):
                    valid_path = full_path
                    break
            if valid_path:
                img = Image.open(valid_path)
                if img.mode != 'RGB': img = img.convert('RGB')
                loaded_images.append(img)
                w, h = img.size
                max_w, max_h = max(max_w, w), max(max_h, h)

        if not loaded_images: return

        MARGIN = 50
        canvas_w = max_w * COLS + PADDING * (COLS - 1) + 2 * MARGIN
        canvas_h = max_h * ROWS + PADDING * (ROWS - 1) + 2 * MARGIN
        combined = Image.new('RGB', (canvas_w, canvas_h), (255, 255, 255))
        draw = ImageDraw.Draw(combined)

        try: font = ImageFont.truetype("arial.ttf", LABEL_SIZE)
        except: font = ImageFont.load_default()

        labels = "abcdef"
        for idx, img in enumerate(loaded_images[:ROWS*COLS]):
            row, col = idx // COLS, idx % COLS
            cell_x = MARGIN + col * (max_w + PADDING)
            cell_y = MARGIN + row * (max_h + PADDING)
            img_w, img_h = img.size
            paste_x = cell_x + (max_w - img_w) // 2
            paste_y = cell_y + (max_h - img_h) // 2
            combined.paste(img, (paste_x, paste_y))
            draw.text((cell_x + LABEL_OFFSET, cell_y + LABEL_OFFSET), labels[idx], fill="black", font=font)

        png_path = os.path.join(OUT_DIR, f"{OUTPUT_NAME}.png")
        tif_path = os.path.join(OUT_DIR, f"{OUTPUT_NAME}.tif")
        combined.save(png_path, dpi=(300, 300))
        combined.save(tif_path, compression="tiff_lzw", dpi=(300, 300))

    combine_images_lowercase()

    ev = epi_starts[epi_starts["episode_no"] >= 2].copy()
    ev = ev.join(t0, on="ID_clean")
    ev["t_event_w"] = ((ev["episode_start"] - ev["t0"]).dt.days / 7.0).astype("float32")
    ev = ev.dropna(subset=["t_event_w"])
    ev = ev[(ev["t_event_w"] >= 0) & (ev["t_event_w"] <= MAX_WEEKS)]
    ev = ev.sort_values(["ID_clean","t_event_w"])

    censor = ((study_end - t0).dt.days / 7.0).astype("float32")
    censor = np.minimum(censor, MAX_WEEKS).rename("censor_w")

    ev["start_w"] = ev.groupby("ID_clean")["t_event_w"].shift(1).fillna(0.0).astype("float32")
    ev["stop_w"] = ev["t_event_w"].astype("float32")
    ev_rows = ev.loc[:, ["ID_clean","start_w","stop_w"]].copy()
    ev_rows["event"] = 1

    last_event = ev.groupby("ID_clean")["t_event_w"].max().rename("last_event_w")
    tmp = pd.DataFrame(index=t0.index)
    tmp = tmp.join(censor).join(last_event)
    tmp["start_w"] = tmp["last_event_w"].fillna(0.0).astype("float32")
    tmp["stop_w"] = tmp["censor_w"].astype("float32")
    tmp["event"] = 0
    tmp = tmp[tmp["stop_w"] > tmp["start_w"]]
    cens_rows = tmp.reset_index().rename(columns={"index":"ID_clean"})[["ID_clean","start_w","stop_w","event"]]

    ag_full = pd.concat([ev_rows, cens_rows], ignore_index=True)
    ag_full = ag_full.merge(baseline.reset_index(), on="ID_clean", how="left")

    exclude = ["Students", "Preschool-aged children"]
    ag_employed = ag_full[~ag_full["Occupation"].isin(exclude)].copy()

    for dset in [ag_full, ag_employed]:
        for c in ["Sex_clean","Age_group","Occupation","City","Case_classification"]:
            dset[c] = dset[c].astype("category")
        dset["event"] = dset["event"].astype("int8")

    ag_full.to_csv(os.path.join(OUT_DIR, "AG_long_data_Full.csv"), index=False)
    ag_employed.to_csv(os.path.join(OUT_DIR, "AG_long_data_Employed.csv"), index=False)

    CONTROL_SAMPLE_RATIO = 10
    RANDOM_SEED = 2026

    REF = {
        "Sex_clean": "Female",
        "Age_group": "20-45",
        "Occupation": "Farmers",
        "City": "Rural",
        "Case_classification": "Clinical diagnosis cases",
    }
    COVARS = ["Sex_clean", "Age_group", "Occupation", "City", "Case_classification"]

    def build_dummies(df_d, col, ref_value):
        s = df_d[col].astype(str)
        d = pd.get_dummies(s, prefix=col)
        ref_col = f"{col}_{ref_value}"
        if ref_col in d.columns: d = d.drop(columns=[ref_col])
        return d.astype("int8")

    def make_X(df_x, covars, ref_map):
        return pd.concat([build_dummies(df_x, c, ref_map[c]) for c in covars], axis=1)

    baseline_emp = baseline[~baseline["Occupation"].isin(exclude)].copy()
    emp_ids = baseline_emp.index

    ag_emp = ag_full[ag_full["ID_clean"].isin(emp_ids)].copy()
    t0_emp = t0.loc[emp_ids]
    event_date_emp = event_date.loc[emp_ids]

    ag2 = ag_emp.copy()
    ag2["pt"] = (ag2["stop_w"] - ag2["start_w"]).astype("float32")
    ag2 = ag2[ag2["pt"] > 0]

    evt_df = ag2[ag2["event"] == 1].copy()
    non_sample = ag2[ag2["event"] == 0].sample(frac=1.0/CONTROL_SAMPLE_RATIO, random_state=RANDOM_SEED).copy()
    ag_s = pd.concat([evt_df, non_sample], ignore_index=True)
    ag_s["w"] = np.where(ag_s["event"] == 1, 1.0, float(CONTROL_SAMPLE_RATIO)).astype("float32")

    X_ag = make_X(ag_s, COVARS, REF)
    res_multi = sm.GLM(ag_s["event"], sm.add_constant(X_ag), family=sm.families.Poisson(), 
                       offset=np.log(ag_s["pt"]), var_weights=ag_s["w"]).fit(cov_type="cluster", cov_kwds={"groups": ag_s["ID_clean"]})

    multi_ag = pd.DataFrame([{"term": t, "aHR": np.exp(res_multi.params[t]), "p": res_multi.pvalues[t]} for t in X_ag.columns])
    multi_ag.to_csv(os.path.join(OUT_DIR, "Employed_Only_AG_Results.csv"), index=False)

    c_w = ((study_end - t0_emp).dt.days / 7.0).astype("float32")
    censor_w = np.minimum(c_w, MAX_WEEKS)
    e_w = ((event_date_emp - t0_emp).dt.days / 7.0).astype("float32")
    e_flag = (e_w.notna() & (e_w <= MAX_WEEKS) & (e_w >= 0))

    cox_df = baseline_emp.copy()
    cox_df["duration"] = np.where(e_flag, e_w, censor_w).astype("float32")
    cox_df["event"] = e_flag.astype("int8")
    cox_df = cox_df[cox_df["duration"] > 0]

    X_cox = make_X(cox_df.reset_index(), COVARS, REF)
    cox_model_df = pd.concat([cox_df[["duration", "event"]].reset_index(drop=True), X_cox.reset_index(drop=True)], axis=1)

    cph = CoxPHFitter()
    cph.fit(cox_model_df, duration_col="duration", event_col="event", robust=True)
    cph.summary.to_csv(os.path.join(OUT_DIR, "Cox_Employed_Summary.csv"))

    def load_data(path):
        for seps in ['\t', ',']:
            for enc in ['utf-8', 'utf-8-sig', 'latin1']:
                try:
                    temp_df = pd.read_csv(path, sep=seps, encoding=enc, quoting=3)
                    if temp_df.shape[1] > 3:
                        return temp_df
                except:
                    continue
        return None

    df_kde = load_data(KDE_INPUT_PATH)

    if df_kde is not None:
        col_map_kde = {}
        for c in df_kde.columns:
            c_clean = str(c).strip().lower()
            if any(x in c_clean for x in ['id']): col_map_kde['id'] = c
            if 'age' in c_clean: col_map_kde['age'] = c
            if 'date' in c_clean: col_map_kde['date'] = c

        df_kde[col_map_kde['id']] = df_kde[col_map_kde['id']].astype(str).str.replace("'", "").str.strip()
        df_kde['Clean_Date'] = pd.to_datetime(df_kde[col_map_kde['date']], errors='coerce')

        def extract_age_kde(age_val):
            if pd.isna(age_val): return None
            nums = re.findall(r'\d+', str(age_val))
            return int(nums[0]) if nums else None

        df_kde['Age_Numeric'] = df_kde[col_map_kde['age']].apply(extract_age_kde)
        df_kde = df_kde.dropna(subset=[col_map_kde['id'], 'Clean_Date', 'Age_Numeric'])
        df_kde = df_kde.sort_values(by=[col_map_kde['id'], 'Clean_Date'])
        df_kde['interval_days'] = df_kde.groupby(col_map_kde['id'])['Clean_Date'].diff().dt.days

        reinfected_df = df_kde[df_kde['interval_days'] > 14].copy()
        reinfected_df['interval_months'] = reinfected_df['interval_days'] / 30.44

        def age_group_precise(age):
            if age <= 5: return '0-5 yrs'
            elif 6 <= age <= 60: return '6-60 yrs'
            elif age >= 61: return '>=61 yrs'
            else: return None

        reinfected_df['Age Group'] = reinfected_df['Age_Numeric'].apply(age_group_precise)
        plot_df = reinfected_df[reinfected_df['interval_months'] <= 36].copy()

        plt.style.use('seaborn-v0_8-white')
        plt.rcParams['font.sans-serif'] = ['Arial']
        fig, ax = plt.subplots(figsize=(12, 6))

        colors = {"0-5 yrs": "#084594", "6-60 yrs": "#525252", ">=61 yrs": "#CB181D"}
        h_order = ["0-5 yrs", "6-60 yrs", ">=61 yrs"]

        sns.kdeplot(data=plot_df, x='interval_months', hue='Age Group', 
                    hue_order=h_order, palette=colors,
                    fill=False, common_norm=False, linewidth=3, 
                    bw_adjust=0.8, 
                    clip=(0, 30), 
                    ax=ax)

        plt.xlabel('Months Since Previous Infection', fontsize=12)
        plt.ylabel('Density (Kernel Density Estimation)', fontsize=12)
        plt.xlim(0, 30)
        plt.xticks(range(0, 30, 3))
        plt.grid(axis='y', linestyle='-', alpha=0.1)
        sns.despine()

        from matplotlib.lines import Line2D
        custom_legend = [Line2D([0], [0], color=colors[g], lw=3, label=g) for g in h_order]
        ax.legend(handles=custom_legend, title='Age Group', loc='upper right')

        plt.savefig(os.path.join(OUT_DIR, 'Fig_Final_36Months.png'), dpi=300, bbox_inches='tight')

