# Copyright (c) 2026 INNOVATIO SAS
# SPDX-License-Identifier: MIT
"""Reproduce the synthetic MSC-P-030 retention-survival example.

Standard-library only. The script generates a deterministic teaching dataset,
then computes Kaplan-Meier estimates, a log-rank comparison, a Cox model and a
declared time-interaction diagnostic for proportional hazards.
"""

from __future__ import annotations

import argparse
import csv
import math
import random
from datetime import date, timedelta
from pathlib import Path

SEED = 20260820
N = 600


def generate(path: Path) -> None:
    rng = random.Random(SEED)
    rows = []
    for index in range(1, N + 1):
        origin = date(2024, 1, 1) + timedelta(days=(index - 1) % 90)
        engagement = rng.gauss(0.0, 1.0)
        annual_probability = 1.0 / (1.0 + math.exp(-(-0.25 + 0.75 * engagement)))
        annual = int(rng.random() < annual_probability)
        monthly_hazard = math.exp(-2.58 - 0.52 * annual - 0.24 * engagement)
        churn_time = -math.log(max(rng.random(), 1e-12)) / monthly_hazard
        if rng.random() < 0.72:
            censor_time = 24.0
            censor_reason = "administrative_end"
        else:
            censor_time = 8.0 + 16.0 * rng.random()
            censor_reason = "observation_end"
        event = int(churn_time <= censor_time)
        duration = min(churn_time, censor_time)
        end_date = origin + timedelta(days=round(duration * 30.4375))
        rows.append(
            {
                "customer_id": f"C{index:04d}",
                "plan_group": "annual" if annual else "monthly",
                "baseline_engagement_z": f"{engagement:.6f}",
                "start_date": origin.isoformat(),
                "end_date": end_date.isoformat(),
                "duration_months": f"{duration:.6f}",
                "churn_event": str(event),
                "censor_reason": "not_censored" if event else censor_reason,
            }
        )
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("w", encoding="utf-8", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)


def load(path: Path) -> list[dict[str, float]]:
    rows = []
    invalid = 0
    required = ("duration_months", "churn_event", "plan_group", "baseline_engagement_z")
    with path.open(encoding="utf-8", newline="") as handle:
        for row in csv.DictReader(handle):
            try:
                if any(not row.get(field, "").strip() for field in required):
                    raise ValueError("missing required value")
                duration = float(row["duration_months"])
                event = int(row["churn_event"])
                engagement = float(row["baseline_engagement_z"])
                if not math.isfinite(duration) or duration <= 0 or event not in (0, 1) or row["plan_group"] not in ("monthly", "annual") or not math.isfinite(engagement):
                    raise ValueError("invalid required value")
            except (TypeError, ValueError):
                invalid += 1
                continue
            rows.append(
                {
                    "duration": duration,
                    "event": float(event),
                    "annual": float(row["plan_group"] == "annual"),
                    "engagement": engagement,
                }
            )
    if invalid:
        raise ValueError(f"{invalid} invalid or incomplete rows; analysis aborted without imputation")
    if not rows:
        raise ValueError("dataset is empty")
    return rows


def km(rows: list[dict[str, float]], horizon: float) -> tuple[float, float, float]:
    survival = 1.0
    greenwood = 0.0
    for time in sorted({row["duration"] for row in rows if row["event"] == 1 and row["duration"] <= horizon}):
        at_risk = sum(row["duration"] >= time for row in rows)
        events = sum(row["event"] == 1 and row["duration"] == time for row in rows)
        survival *= 1.0 - events / at_risk
        if at_risk > events:
            greenwood += events / (at_risk * (at_risk - events))
    if not 0.0 < survival < 1.0 or greenwood == 0.0:
        return survival, float("nan"), float("nan")
    log_survival = math.log(survival)
    se_loglog = math.sqrt(greenwood) / abs(log_survival)
    center = math.log(-log_survival)
    lower = math.exp(-math.exp(center + 1.959964 * se_loglog))
    upper = math.exp(-math.exp(center - 1.959964 * se_loglog))
    return survival, lower, upper


def median_survival(rows: list[dict[str, float]]) -> float | None:
    survival = 1.0
    for time in sorted({row["duration"] for row in rows if row["event"] == 1}):
        at_risk = sum(row["duration"] >= time for row in rows)
        events = sum(row["event"] == 1 and row["duration"] == time for row in rows)
        survival *= 1.0 - events / at_risk
        if survival <= 0.5:
            return time
    return None


def logrank(rows: list[dict[str, float]]) -> tuple[float, float]:
    observed_minus_expected = 0.0
    variance = 0.0
    for time in sorted({row["duration"] for row in rows if row["event"] == 1}):
        risk = [row for row in rows if row["duration"] >= time]
        events = [row for row in rows if row["event"] == 1 and row["duration"] == time]
        n, n1, d, d1 = len(risk), sum(row["annual"] for row in risk), len(events), sum(row["annual"] for row in events)
        observed_minus_expected += d1 - d * n1 / n
        if n > 1:
            variance += d * (n1 / n) * (1.0 - n1 / n) * (n - d) / (n - 1)
    if variance <= 0.0:
        raise ValueError("log-rank variance is zero")
    chi2 = observed_minus_expected**2 / variance
    return chi2, math.erfc(math.sqrt(chi2 / 2.0))


def solve(matrix: list[list[float]], vector: list[float]) -> list[float]:
    n = len(vector)
    augmented = [row[:] + [value] for row, value in zip(matrix, vector)]
    for column in range(n):
        pivot = max(range(column, n), key=lambda row: abs(augmented[row][column]))
        if abs(augmented[pivot][column]) < 1e-12:
            raise ValueError("singular information matrix")
        augmented[column], augmented[pivot] = augmented[pivot], augmented[column]
        scale = augmented[column][column]
        augmented[column] = [value / scale for value in augmented[column]]
        for row in range(n):
            if row == column:
                continue
            factor = augmented[row][column]
            augmented[row] = [a - factor * b for a, b in zip(augmented[row], augmented[column])]
    return [augmented[row][-1] for row in range(n)]


def inverse(matrix: list[list[float]]) -> list[list[float]]:
    n = len(matrix)
    return [[solve(matrix, [float(i == column) for i in range(n)])[row] for column in range(n)] for row in range(n)]


def cox(rows: list[dict[str, float]], time_interaction: bool = False) -> tuple[list[float], list[list[float]], int]:
    size = 3 if time_interaction else 2
    beta = [0.0] * size
    event_times = sorted({row["duration"] for row in rows if row["event"] == 1})
    information = [[0.0] * size for _ in range(size)]
    for iteration in range(1, 51):
        score = [0.0] * size
        information = [[0.0] * size for _ in range(size)]
        for time in event_times:
            events = [row for row in rows if row["event"] == 1 and row["duration"] == time]
            risk = [row for row in rows if row["duration"] >= time]
            def values(row: dict[str, float]) -> list[float]:
                base = [row["annual"], row["engagement"]]
                return base + ([row["annual"] * math.log(time / 12.0)] if time_interaction else [])
            weighted = []
            for row in risk:
                x = values(row)
                weight = math.exp(sum(coef * value for coef, value in zip(beta, x)))
                weighted.append((weight, x))
            s0 = sum(weight for weight, _ in weighted)
            means = [sum(weight * x[j] for weight, x in weighted) / s0 for j in range(size)]
            second = [[sum(weight * x[j] * x[k] for weight, x in weighted) / s0 for k in range(size)] for j in range(size)]
            for event in events:
                x_event = values(event)
                for j in range(size):
                    score[j] += x_event[j] - means[j]
            d = len(events)
            for j in range(size):
                for k in range(size):
                    information[j][k] += d * (second[j][k] - means[j] * means[k])
        step = solve(information, score)
        beta = [value + increment for value, increment in zip(beta, step)]
        if max(abs(value) for value in step) < 1e-10:
            return beta, inverse(information), iteration
    raise RuntimeError("Cox model did not converge in 50 iterations")


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--csv", type=Path, default=Path(__file__).with_name("msc-p030-retention-survival.csv"))
    parser.add_argument("--generate", action="store_true")
    args = parser.parse_args()
    if args.generate:
        generate(args.csv)
    rows = load(args.csv)
    print(f"n={len(rows)} events={sum(row['event'] for row in rows):.0f} censored={sum(1-row['event'] for row in rows):.0f}")
    for group_name, group_value in (("monthly", 0.0), ("annual", 1.0)):
        group = [row for row in rows if row["annual"] == group_value]
        s9, low9, high9 = km(group, 9.0)
        s12, low12, high12 = km(group, 12.0)
        s18, low18, high18 = km(group, 18.0)
        median = median_survival(group)
        print(f"{group_name}: n={len(group)} S9={s9:.6f} CI95=[{low9:.6f},{high9:.6f}] S12={s12:.6f} CI95=[{low12:.6f},{high12:.6f}] S18={s18:.6f} CI95=[{low18:.6f},{high18:.6f}] median={median if median is not None else 'not_reached'}")
    chi2, p_value = logrank(rows)
    print(f"logrank: chi2={chi2:.6f} p={p_value:.6e}")
    beta, covariance, iterations = cox(rows)
    se = math.sqrt(covariance[0][0])
    print(f"cox_adjusted: beta_annual={beta[0]:.6f} HR={math.exp(beta[0]):.6f} CI95=[{math.exp(beta[0]-1.959964*se):.6f},{math.exp(beta[0]+1.959964*se):.6f}] beta_engagement={beta[1]:.6f} iterations={iterations}")
    beta_ph, covariance_ph, iterations_ph = cox(rows, time_interaction=True)
    gamma, gamma_se = beta_ph[2], math.sqrt(covariance_ph[2][2])
    z = gamma / gamma_se
    p_ph = math.erfc(abs(z) / math.sqrt(2.0))
    print(f"ph_time_interaction: gamma={gamma:.6f} SE={gamma_se:.6f} z={z:.6f} p={p_ph:.8f} iterations={iterations_ph}")
    print("scope: synthetic observational association; no causal or individual-probability claim")


if __name__ == "__main__":
    main()
