# Copyright (c) 2026 INNOVATIO SAS
# SPDX-License-Identifier: MIT
"""MSC-P-043 (marketing-science-center.com): project subscriber retention and value with the sBG model.

Standard library only. Synthetic cohort, declared as such: 2,000 monthly subscribers acquired
together, whose individual churn probabilities follow a Beta(0.5, 2.5) distribution (mean monthly
churn 1/6). Each lifetime is drawn by inversion of the shifted-beta-geometric survivor function.
Six renewal decisions are used for calibration; months 7 to 24 are kept as a holdout. Two models
are fitted by maximum likelihood: a geometric model (one constant retention rate for everyone) and
the sBG model (constant individual rate, beta heterogeneity across subscribers). Values use a
monthly margin of 12 booked at the start of each paid month and a monthly discount rate of 1%.
The R reference prints exactly the same lines.
"""
import math

LCG_SEED = 20261002
LCG_MULTIPLIER = 1103515245
LCG_INCREMENT = 12345
LCG_MODULUS = 2147483648

SUBSCRIBERS = 2000
TRUE_ALPHA = 0.5
TRUE_BETA = 2.5
CALIBRATION = 6
HORIZON = 24
MARGIN = 12.0
DISCOUNT = 0.01
TERMS = 3000
BOOTSTRAP = 200
REPLICATIONS = 100
NM_MAX_ITER = 5000
NM_TOL = 1e-12
HESSIAN_STEP = 1e-4


class Declared:
    def __init__(self, seed):
        self.state = seed

    def uniform(self):
        self.state = (LCG_MULTIPLIER * self.state + LCG_INCREMENT) % LCG_MODULUS
        return (self.state + 0.5) / LCG_MODULUS


def sbg_survival(alpha, beta, last):
    """S(0..last) by the forward recursion r_t = (beta + t - 1) / (alpha + beta + t - 1)."""
    s = [1.0]
    for t in range(1, last + 1):
        s.append(s[-1] * (beta + t - 1.0) / (alpha + beta + t - 1.0))
    return s


def draw_lifetimes(stream, alpha, beta, count, cap):
    """Lifetime T = min{t >= 1 : u >= S(t)}; values above cap are returned as cap + 1."""
    s = sbg_survival(alpha, beta, cap)
    out = []
    for _ in range(count):
        u = stream.uniform()
        t = 1
        while t <= cap and u < s[t]:
            t += 1
        out.append(t)
    return out


def calibration_counts(lifetimes, horizon):
    churned = [0] * (horizon + 1)
    for t in lifetimes:
        if t <= horizon:
            churned[t] += 1
    survivors = len(lifetimes) - sum(churned)
    return churned, survivors


def sbg_loglik(alpha, beta, churned, survivors, horizon):
    p = alpha / (alpha + beta)
    s = 1.0 - p
    ll = churned[1] * math.log(p) if churned[1] > 0 else 0.0
    for t in range(2, horizon + 1):
        p = p * (beta + t - 2.0) / (alpha + beta + t - 1.0)
        s = s - p
        if churned[t] > 0:
            ll += churned[t] * math.log(p)
    ll += survivors * math.log(s)
    return ll


def nelder_mead(f, start, step):
    pts = [list(start), [start[0] + step, start[1]], [start[0], start[1] + step]]
    vals = [f(p) for p in pts]
    converged = False
    for _ in range(NM_MAX_ITER):
        order = sorted(range(3), key=lambda i: (vals[i], i))
        pts = [pts[i] for i in order]
        vals = [vals[i] for i in order]
        size = max(abs(pts[i][k] - pts[0][k]) for i in (1, 2) for k in (0, 1))
        if vals[2] - vals[0] < NM_TOL and size < 1e-9:
            converged = True
            break
        centroid = [(pts[0][k] + pts[1][k]) / 2.0 for k in (0, 1)]
        refl = [centroid[k] + (centroid[k] - pts[2][k]) for k in (0, 1)]
        fr = f(refl)
        if fr < vals[0]:
            exp_pt = [centroid[k] + 2.0 * (centroid[k] - pts[2][k]) for k in (0, 1)]
            fe = f(exp_pt)
            if fe < fr:
                pts[2], vals[2] = exp_pt, fe
            else:
                pts[2], vals[2] = refl, fr
        elif fr < vals[1]:
            pts[2], vals[2] = refl, fr
        else:
            if fr < vals[2]:
                con = [centroid[k] + 0.5 * (refl[k] - centroid[k]) for k in (0, 1)]
            else:
                con = [centroid[k] + 0.5 * (pts[2][k] - centroid[k]) for k in (0, 1)]
            fc = f(con)
            if fc < min(fr, vals[2]):
                pts[2], vals[2] = con, fc
            else:
                for i in (1, 2):
                    pts[i] = [pts[0][k] + 0.5 * (pts[i][k] - pts[0][k]) for k in (0, 1)]
                    vals[i] = f(pts[i])
    best = min(range(3), key=lambda i: (vals[i], i))
    return pts[best], vals[best], converged


def fit_sbg(churned, survivors, horizon):
    def neg(x):
        return -sbg_loglik(math.exp(x[0]), math.exp(x[1]), churned, survivors, horizon)

    x, value, converged = nelder_mead(neg, [0.0, 0.0], 0.5)
    if not converged:
        raise SystemExit("error: the sBG fit did not converge within NM_MAX_ITER iterations")
    return math.exp(x[0]), math.exp(x[1]), -value, x, neg


def fit_geometric(churned, survivors, horizon):
    deaths = sum(churned[1:horizon + 1])
    exposure = sum(t * churned[t] for t in range(1, horizon + 1)) + horizon * survivors
    theta = deaths / exposure
    ll = deaths * math.log(theta) + (exposure - deaths) * math.log(1.0 - theta)
    return theta, ll, deaths, exposure


def hessian_se(neg, x):
    h = HESSIAN_STEP
    f0 = neg(x)
    def at(a, b):
        return neg([x[0] + a, x[1] + b])
    h11 = (at(h, 0) - 2.0 * f0 + at(-h, 0)) / (h * h)
    h22 = (at(0, h) - 2.0 * f0 + at(0, -h)) / (h * h)
    h12 = (at(h, h) - at(h, -h) - at(-h, h) + at(-h, -h)) / (4.0 * h * h)
    det = h11 * h22 - h12 * h12
    if not (h11 > 0.0 and det > 0.0):
        raise SystemExit("error: the numerical Hessian is not positive definite at the sBG optimum")
    # Standard errors on the log scale; the delta method gives SE(alpha) = alpha * SE(ln alpha).
    return math.sqrt(h22 / det), math.sqrt(h11 / det)


def clv_new(surv, rate=DISCOUNT, payments=TERMS + 1):
    """Payment t + 1 is booked at the start of month t + 1 with probability S(t); discounted by (1 + rate)^t."""
    total = 0.0
    factor = 1.0
    for t in range(payments):
        total += surv[t] * factor
        factor /= 1.0 + rate
    return MARGIN * total


def residual_value(surv, n, rate=DISCOUNT, last=TERMS):
    """Value at the start of month n + 1 of payments from month n + 2 on, given T > n."""
    total = 0.0
    factor = 1.0
    for t in range(n + 1, last + 1):
        factor /= 1.0 + rate
        total += surv[t] / surv[n] * factor
    return MARGIN * total


def derl(surv, n, rate):
    """Fader and Hardie (2010) eq. (4): sum over t >= n of S(t) / S(n - 1) / (1 + rate)^(t - n)."""
    total = 0.0
    factor = 1.0
    for t in range(n, TERMS + 1):
        total += surv[t] / surv[n - 1] * factor
        factor /= 1.0 + rate
    return total


PUBLISHED = {
    # Fader and Hardie (2007), Appendix B (alpha 0.668; p. 9 prints 0.688) and Section 3.
    "high_end": "alpha=0.668 beta=3.806 loglik=-1.611",
    "regular": "alpha=0.704 beta=1.182",
    # Fader and Hardie (2010), Table 4, DERL column, n = 5 down to 1, d = 10 %.
    "case1": "3.84 3.72 3.59 3.45 3.31",
    "case2": "10.19 10.06 9.86 9.46 7.68",
}


def published_checks():
    """Reproduce values printed in the two sources before trusting the code on new data; stop if one is missed."""
    for name, s in (("high_end", [1.0, 0.869, 0.743, 0.653, 0.593, 0.551, 0.517, 0.491]),
                    ("regular", [1.0, 0.631, 0.468, 0.382, 0.326, 0.289, 0.262, 0.241])):
        churned = [0.0] + [s[t - 1] - s[t] for t in range(1, 8)]
        a, b, ll, _, _ = fit_sbg(churned, s[7], 7)
        line = f"alpha={fmt(a, 3)} beta={fmt(b, 3)} loglik={fmt(ll, 3)}"
        if not line.startswith(PUBLISHED[name]):
            raise SystemExit(f"error: Fader and Hardie (2007) {name} not reproduced: {line}")
        print(f"metric.check.fader_hardie_2007_{name} {line}")
    # Case 2 is defined by mean 0.20 and polarization 0.75, i.e. alpha = 1/15 and beta = 4/15 (printed 0.067, 0.267).
    for name, a, b in (("case1", 3.8, 15.2), ("case2", 1.0 / 15.0, 4.0 / 15.0)):
        surv = sbg_survival(a, b, TERMS)
        line = " ".join(fmt(derl(surv, n, 0.1), 2) for n in (5, 4, 3, 2, 1))
        if line != PUBLISHED[name]:
            raise SystemExit(f"error: Fader and Hardie (2010) Table 4 {name} not reproduced: {line}")
        print(f"metric.check.fader_hardie_2010_{name}_derl_n5_to_n1={line}")


def geometric_survival(theta, last):
    return [(1.0 - theta) ** t for t in range(last + 1)]


def plain_sum(values):
    """Left-to-right double sum; Python 3.12+ sum() compensates, R's sum() uses long double."""
    total = 0.0
    for v in values:
        total += v
    return total


def fmt(x, digits):
    text = f"{x:.{digits}f}"
    return text[1:] if text.startswith("-") and float(text) == 0.0 else text


def summary(values, truth):
    n = len(values)
    mean = plain_sum(values) / n
    sd = math.sqrt(plain_sum([(v - mean) * (v - mean) for v in values]) / (n - 1))
    rmse = math.sqrt(plain_sum([(v - truth) * (v - truth) for v in values]) / n)
    return mean, sd, mean - truth, rmse


def main():
    stream = Declared(LCG_SEED)
    true_surv = sbg_survival(TRUE_ALPHA, TRUE_BETA, TERMS)
    true_clv = clv_new(true_surv)
    true_rv = residual_value(true_surv, CALIBRATION)

    print(f"metric.design.subscribers={SUBSCRIBERS} calibration_months={CALIBRATION} holdout_to_month={HORIZON}")
    print(f"metric.design.seed={LCG_SEED} bootstrap={BOOTSTRAP} replications={REPLICATIONS}")
    print(f"metric.design.true_alpha={fmt(TRUE_ALPHA, 2)} true_beta={fmt(TRUE_BETA, 2)} "
          f"true_mean_churn={fmt(TRUE_ALPHA / (TRUE_ALPHA + TRUE_BETA), 4)} "
          f"true_polarization={fmt(1.0 / (TRUE_ALPHA + TRUE_BETA + 1.0), 4)}")
    print(f"metric.design.margin_per_month={fmt(MARGIN, 2)} discount_per_month={fmt(DISCOUNT, 4)} series_terms={TERMS}")
    print(f"metric.truth.clv_new_subscriber={fmt(true_clv, 2)} residual_value_per_survivor={fmt(true_rv, 2)} "
          f"s24={fmt(true_surv[HORIZON], 4)}")
    published_checks()

    lifetimes = draw_lifetimes(stream, TRUE_ALPHA, TRUE_BETA, SUBSCRIBERS, HORIZON)
    churned, survivors = calibration_counts(lifetimes, CALIBRATION)
    active = [SUBSCRIBERS]
    for t in range(1, HORIZON + 1):
        active.append(sum(1 for v in lifetimes if v > t))
    print("metric.data.active_by_month=" + " ".join(f"{t}:{active[t]}" for t in range(0, HORIZON + 1)))
    print("metric.data.retention_rate_calibration=" + " ".join(
        f"{t}:{fmt(active[t] / active[t - 1], 4)}" for t in range(1, CALIBRATION + 1)))
    print("metric.data.retention_rate_holdout=" + " ".join(
        f"{t}:{fmt(active[t] / active[t - 1], 4)}" for t in range(CALIBRATION + 1, HORIZON + 1)))

    theta, ll_geo, deaths, exposure = fit_geometric(churned, survivors, CALIBRATION)
    alpha, beta, ll_sbg, x, neg = fit_sbg(churned, survivors, CALIBRATION)
    se_log_a, se_log_b = hessian_se(neg, x)
    print(f"metric.fit.geometric theta={fmt(theta, 4)} retention={fmt(1.0 - theta, 4)} "
          f"churners={deaths} renewal_decisions={exposure} loglik={fmt(ll_geo, 3)}")
    print(f"metric.fit.sbg alpha={fmt(alpha, 4)} se={fmt(alpha * se_log_a, 4)} beta={fmt(beta, 4)} "
          f"se={fmt(beta * se_log_b, 4)} loglik={fmt(ll_sbg, 3)}")
    print(f"metric.fit.sbg mean_churn={fmt(alpha / (alpha + beta), 4)} "
          f"polarization={fmt(1.0 / (alpha + beta + 1.0), 4)}")
    print(f"metric.fit.likelihood_ratio={fmt(2.0 * (ll_sbg - ll_geo), 3)} aic_geometric={fmt(-2.0 * ll_geo + 2.0, 3)} "
          f"aic_sbg={fmt(-2.0 * ll_sbg + 4.0, 3)}")

    last_theta = 1.0 - active[CALIBRATION] / active[CALIBRATION - 1]
    print(f"metric.fit.last_rate theta={fmt(last_theta, 4)} retention={fmt(1.0 - last_theta, 4)}")
    s_sbg = sbg_survival(alpha, beta, TERMS)
    s_geo = geometric_survival(theta, TERMS)
    s_last = [active[t] / SUBSCRIBERS if t <= CALIBRATION else
              active[CALIBRATION] / SUBSCRIBERS * (1.0 - last_theta) ** (t - CALIBRATION) for t in range(TERMS + 1)]
    for t in (6, 12, 18, 24):
        print(f"metric.projection.month={t} observed={fmt(active[t] / SUBSCRIBERS, 4)} "
              f"geometric={fmt(s_geo[t], 4)} last_rate={fmt(s_last[t], 4)} sbg={fmt(s_sbg[t], 4)} "
              f"truth={fmt(true_surv[t], 4)}")
    mae_geo = plain_sum([abs(s_geo[t] - active[t] / SUBSCRIBERS) for t in range(CALIBRATION + 1, HORIZON + 1)]) / (HORIZON - CALIBRATION)
    mae_sbg = plain_sum([abs(s_sbg[t] - active[t] / SUBSCRIBERS) for t in range(CALIBRATION + 1, HORIZON + 1)]) / (HORIZON - CALIBRATION)
    mae_last = plain_sum([abs(s_last[t] - active[t] / SUBSCRIBERS) for t in range(CALIBRATION + 1, HORIZON + 1)]) / (HORIZON - CALIBRATION)
    print(f"metric.projection.holdout_mean_abs_error_points geometric={fmt(100.0 * mae_geo, 2)} "
          f"last_rate={fmt(100.0 * mae_last, 2)} sbg={fmt(100.0 * mae_sbg, 2)}")
    print("metric.projection.sbg_retention_rate=" + " ".join(
        f"{t}:{fmt(s_sbg[t] / s_sbg[t - 1], 4)}" for t in (1, 6, 12, 24)))

    clv_geo = clv_new(s_geo)
    clv_sbg = clv_new(s_sbg)
    rv_geo = residual_value(s_geo, CALIBRATION)
    rv_sbg = residual_value(s_sbg, CALIBRATION)
    rv_last = residual_value(geometric_survival(last_theta, TERMS), CALIBRATION)
    print(f"metric.value.clv_new_subscriber geometric={fmt(clv_geo, 2)} sbg={fmt(clv_sbg, 2)} truth={fmt(true_clv, 2)}")
    print(f"metric.value.residual_per_survivor geometric={fmt(rv_geo, 2)} last_rate={fmt(rv_last, 2)} "
          f"sbg={fmt(rv_sbg, 2)} truth={fmt(true_rv, 2)}")
    print(f"metric.value.residual_cohort survivors={survivors} geometric={fmt(survivors * rv_geo, 0)} "
          f"last_rate={fmt(survivors * rv_last, 0)} sbg={fmt(survivors * rv_sbg, 0)} truth={fmt(survivors * true_rv, 0)}")
    print(f"metric.value.gap_vs_truth clv_new_geometric={fmt(100.0 * (clv_geo / true_clv - 1.0), 1)}% "
          f"clv_new_sbg={fmt(100.0 * (clv_sbg / true_clv - 1.0), 1)}% "
          f"residual_geometric={fmt(100.0 * (rv_geo / true_rv - 1.0), 1)}% "
          f"residual_last_rate={fmt(100.0 * (rv_last / true_rv - 1.0), 1)}% "
          f"residual_sbg={fmt(100.0 * (rv_sbg / true_rv - 1.0), 1)}%")

    # How much of each value lies beyond month h. Index t pays in month t + 1, so months 1..h are t <= h - 1.
    print("metric.horizon.share_of_value_beyond_month=" + " ".join(
        f"{h}:clv_sbg={fmt(100.0 * (1.0 - clv_new(s_sbg, payments=h) / clv_sbg), 1)}%"
        f",residual_sbg={fmt(100.0 * (1.0 - residual_value(s_sbg, CALIBRATION, last=h - 1) / rv_sbg), 1)}%"
        for h in (24, 60, 120)))
    print("metric.horizon.clv_first_payments=" + " ".join(
        f"{p}:geometric={fmt(clv_new(s_geo, payments=p), 2)},sbg={fmt(clv_new(s_sbg, payments=p), 2)}"
        f",truth={fmt(clv_new(true_surv, payments=p), 2)}" for p in (24, 36, 60)))

    # The discount rate per period drives both the values and the size of the constant-rate gap.
    print(f"metric.sensitivity.annual_equivalent_of_monthly_discount={fmt(100.0 * ((1.0 + DISCOUNT) ** 12 - 1.0), 1)}%")
    s_last_geo = geometric_survival(last_theta, TERMS)
    for rate in (0.005, 0.01, 0.015, 0.02, 0.10):
        t_clv = clv_new(true_surv, rate)
        t_rv = residual_value(true_surv, CALIBRATION, rate)
        g_rv = residual_value(s_geo, CALIBRATION, rate)
        l_rv = residual_value(s_last_geo, CALIBRATION, rate)
        print(f"metric.sensitivity.discount={fmt(100.0 * rate, 1)}% clv_geometric={fmt(clv_new(s_geo, rate), 2)} "
              f"clv_sbg={fmt(clv_new(s_sbg, rate), 2)} clv_truth={fmt(t_clv, 2)} "
              f"residual_sbg={fmt(residual_value(s_sbg, CALIBRATION, rate), 2)} residual_truth={fmt(t_rv, 2)} "
              f"residual_gap_geometric={fmt(100.0 * (g_rv / t_rv - 1.0), 1)}% "
              f"residual_gap_last_rate={fmt(100.0 * (l_rv / t_rv - 1.0), 1)}%")

    boot_clv = []
    boot_rv = []
    for _ in range(BOOTSTRAP):
        b_life = draw_lifetimes(stream, alpha, beta, SUBSCRIBERS, CALIBRATION)
        b_churned, b_surv = calibration_counts(b_life, CALIBRATION)
        ba, bb, _, _, _ = fit_sbg(b_churned, b_surv, CALIBRATION)
        bs = sbg_survival(ba, bb, TERMS)
        boot_clv.append(clv_new(bs))
        boot_rv.append(residual_value(bs, CALIBRATION))
    boot_clv.sort()
    boot_rv.sort()
    print(f"metric.uncertainty.parametric_bootstrap clv_new_95={fmt(boot_clv[4], 2)}..{fmt(boot_clv[194], 2)} "
          f"residual_per_survivor_95={fmt(boot_rv[4], 2)}..{fmt(boot_rv[194], 2)}")

    rep = {"clv_geo": [], "clv_sbg": [], "rv_geo": [], "rv_last": [], "rv_sbg": [], "s24_geo": [], "s24_sbg": []}
    geo_below = 0
    last_below = 0
    for _ in range(REPLICATIONS):
        r_life = draw_lifetimes(stream, TRUE_ALPHA, TRUE_BETA, SUBSCRIBERS, CALIBRATION)
        r_churned, r_surv = calibration_counts(r_life, CALIBRATION)
        r_last = r_churned[CALIBRATION] / (r_surv + r_churned[CALIBRATION])
        r_theta, _, _, _ = fit_geometric(r_churned, r_surv, CALIBRATION)
        ra, rb, _, _, _ = fit_sbg(r_churned, r_surv, CALIBRATION)
        rs = sbg_survival(ra, rb, TERMS)
        rg = geometric_survival(r_theta, TERMS)
        rep["clv_geo"].append(clv_new(rg))
        rep["clv_sbg"].append(clv_new(rs))
        rep["rv_geo"].append(residual_value(rg, CALIBRATION))
        rep["rv_last"].append(residual_value(geometric_survival(r_last, TERMS), CALIBRATION))
        rep["rv_sbg"].append(residual_value(rs, CALIBRATION))
        rep["s24_geo"].append(rg[HORIZON])
        rep["s24_sbg"].append(rs[HORIZON])
        if rep["rv_geo"][-1] < true_rv:
            geo_below += 1
        if rep["rv_last"][-1] < true_rv:
            last_below += 1
    for key, truth, digits in (("clv_geo", true_clv, 2), ("clv_sbg", true_clv, 2),
                               ("rv_geo", true_rv, 2), ("rv_last", true_rv, 2), ("rv_sbg", true_rv, 2),
                               ("s24_geo", true_surv[HORIZON], 4), ("s24_sbg", true_surv[HORIZON], 4)):
        mean, sd, bias, rmse = summary(rep[key], truth)
        print(f"metric.monte_carlo.{key} mean={fmt(mean, digits)} sd={fmt(sd, digits)} "
              f"bias={fmt(bias, digits)} bias_mcse={fmt(sd / math.sqrt(REPLICATIONS), digits)} rmse={fmt(rmse, digits)}")
    print(f"metric.monte_carlo.residual_below_truth geometric={geo_below} last_rate={last_below} of {REPLICATIONS}")


if __name__ == "__main__":
    main()
