# 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.
# Base R only. Same declared recursion, design and algorithm as msc-p043-reference.py; the two
# programs must print exactly the same lines. Survivor vectors are 1-based: s[t + 1] holds S(t).
# Floating-point sums are written as explicit left-to-right loops, never with sum(), which
# accumulates in long double.
LCG_SEED <- 20261002
LCG_MODULUS <- 2147483648
LCG_INCREMENT <- 12345
LCG_HIGH <- 16838   # 1103515245 = 16838 * 65536 + 20077, split to keep products exact
LCG_LOW <- 20077

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

lcg_state <- LCG_SEED
uniform <- function() {
  high <- (LCG_HIGH * lcg_state) %% LCG_MODULUS
  lcg_state <<- (high * 65536 + LCG_LOW * lcg_state + LCG_INCREMENT) %% LCG_MODULUS
  (lcg_state + 0.5) / LCG_MODULUS
}

plain_sum <- function(values) {
  total <- 0
  for (v in values) total <- total + v
  total
}

sbg_survival <- function(alpha, beta, last) {
  s <- numeric(last + 1L)
  s[1] <- 1
  for (t in seq_len(last)) s[t + 1] <- s[t] * (beta + t - 1) / (alpha + beta + t - 1)
  s
}

draw_lifetimes <- function(alpha, beta, count, cap) {
  s <- sbg_survival(alpha, beta, cap)
  out <- integer(count)
  for (i in seq_len(count)) {
    u <- uniform()
    t <- 1L
    while (t <= cap && u < s[t + 1]) t <- t + 1L
    out[i] <- t
  }
  out
}

calibration_counts <- function(lifetimes, horizon) {
  churned <- integer(horizon)
  for (t in lifetimes) if (t <= horizon) churned[t] <- churned[t] + 1L
  list(churned = churned, survivors = length(lifetimes) - sum(churned))
}

sbg_loglik <- function(alpha, beta, churned, survivors, horizon) {
  p <- alpha / (alpha + beta)
  s <- 1 - p
  ll <- if (churned[1] > 0) churned[1] * log(p) else 0
  for (t in 2:horizon) {
    p <- p * (beta + t - 2) / (alpha + beta + t - 1)
    s <- s - p
    if (churned[t] > 0) ll <- ll + churned[t] * log(p)
  }
  ll + survivors * log(s)
}

nelder_mead <- function(f, start, step) {
  pts <- list(start, c(start[1] + step, start[2]), c(start[1], start[2] + step))
  vals <- c(f(pts[[1]]), f(pts[[2]]), f(pts[[3]]))
  converged <- FALSE
  for (iter in seq_len(NM_MAX_ITER)) {
    ord <- order(vals, 1:3)
    pts <- pts[ord]
    vals <- vals[ord]
    size <- max(abs(pts[[2]] - pts[[1]]), abs(pts[[3]] - pts[[1]]))
    if (vals[3] - vals[1] < NM_TOL && size < 1e-9) {
      converged <- TRUE
      break
    }
    centroid <- (pts[[1]] + pts[[2]]) / 2
    refl <- centroid + (centroid - pts[[3]])
    fr <- f(refl)
    if (fr < vals[1]) {
      exp_pt <- centroid + 2 * (centroid - pts[[3]])
      fe <- f(exp_pt)
      if (fe < fr) {
        pts[[3]] <- exp_pt; vals[3] <- fe
      } else {
        pts[[3]] <- refl; vals[3] <- fr
      }
    } else if (fr < vals[2]) {
      pts[[3]] <- refl; vals[3] <- fr
    } else {
      con <- if (fr < vals[3]) centroid + 0.5 * (refl - centroid) else centroid + 0.5 * (pts[[3]] - centroid)
      fc <- f(con)
      if (fc < min(fr, vals[3])) {
        pts[[3]] <- con; vals[3] <- fc
      } else {
        for (i in 2:3) {
          pts[[i]] <- pts[[1]] + 0.5 * (pts[[i]] - pts[[1]])
          vals[i] <- f(pts[[i]])
        }
      }
    }
  }
  best <- order(vals, 1:3)[1]
  list(x = pts[[best]], value = vals[best], converged = converged)
}

fail <- function(message) {
  cat(paste0("error: ", message, "\n"), file = stderr())
  quit(status = 1)
}

fit_sbg <- function(churned, survivors, horizon) {
  neg <- function(x) -sbg_loglik(exp(x[1]), exp(x[2]), churned, survivors, horizon)
  res <- nelder_mead(neg, c(0, 0), 0.5)
  if (!res$converged) fail("the sBG fit did not converge within NM_MAX_ITER iterations")
  list(alpha = exp(res$x[1]), beta = exp(res$x[2]), ll = -res$value, x = res$x, neg = neg)
}

fit_geometric <- function(churned, survivors, horizon) {
  deaths <- sum(churned[1:horizon])
  exposure <- sum((1:horizon) * churned[1:horizon]) + horizon * survivors
  theta <- deaths / exposure
  list(theta = theta, ll = deaths * log(theta) + (exposure - deaths) * log(1 - theta),
       deaths = deaths, exposure = exposure)
}

hessian_se <- function(neg, x) {
  h <- HESSIAN_STEP
  f0 <- neg(x)
  at <- function(a, b) neg(c(x[1] + a, x[2] + b))
  h11 <- (at(h, 0) - 2 * f0 + at(-h, 0)) / (h * h)
  h22 <- (at(0, h) - 2 * f0 + at(0, -h)) / (h * h)
  h12 <- (at(h, h) - at(h, -h) - at(-h, h) + at(-h, -h)) / (4 * h * h)
  det <- h11 * h22 - h12 * h12
  if (!(h11 > 0 && det > 0)) fail("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).
  c(sqrt(h22 / det), sqrt(h11 / det))
}

clv_new <- function(surv, rate = DISCOUNT, payments = TERMS + 1L) {
  # Payment t + 1 is booked at the start of month t + 1 with probability S(t); discounted by (1 + rate)^t.
  total <- 0
  factor <- 1
  for (t in 0:(payments - 1L)) {
    total <- total + surv[t + 1] * factor
    factor <- factor / (1 + rate)
  }
  MARGIN * total
}

residual_value <- function(surv, n, rate = DISCOUNT, last = TERMS) {
  total <- 0
  factor <- 1
  for (t in (n + 1):last) {
    factor <- factor / (1 + rate)
    total <- total + surv[t + 1] / surv[n + 1] * factor
  }
  MARGIN * total
}

derl <- function(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
  factor <- 1
  for (t in n:TERMS) {
    total <- total + surv[t + 1] / surv[n] * factor
    factor <- factor / (1 + rate)
  }
  total
}

geometric_survival <- function(theta, last) (1 - theta)^(0:last)

fmt <- function(x, digits) {
  text <- sprintf(paste0("%.", digits, "f"), x)
  if (startsWith(text, "-") && as.numeric(text) == 0) substring(text, 2) else text
}

summary_stats <- function(values, truth) {
  n <- length(values)
  mean <- plain_sum(values) / n
  sd <- sqrt(plain_sum((values - mean) * (values - mean)) / (n - 1))
  rmse <- sqrt(plain_sum((values - truth) * (values - truth)) / n)
  c(mean, sd, mean - truth, rmse)
}

out <- function(...) cat(paste0(..., "\n"), sep = "")

PUBLISHED <- list(
  # 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"
)

published_checks <- function() {
  # Reproduce values printed in the two sources before trusting the code on new data; stop if one is missed.
  data <- list(high_end = c(1, 0.869, 0.743, 0.653, 0.593, 0.551, 0.517, 0.491),
               regular = c(1, 0.631, 0.468, 0.382, 0.326, 0.289, 0.262, 0.241))
  for (name in names(data)) {
    s <- data[[name]]
    churned <- s[1:7] - s[2:8]
    f <- fit_sbg(churned, s[8], 7L)
    line <- paste0("alpha=", fmt(f$alpha, 3), " beta=", fmt(f$beta, 3), " loglik=", fmt(f$ll, 3))
    if (!startsWith(line, PUBLISHED[[name]])) fail(paste0("Fader and Hardie (2007) ", name, " not reproduced: ", line))
    out("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).
  cases <- list(case1 = c(3.8, 15.2), case2 = c(1 / 15, 4 / 15))
  for (name in names(cases)) {
    surv <- sbg_survival(cases[[name]][1], cases[[name]][2], TERMS)
    line <- paste(vapply(5:1, function(n) fmt(derl(surv, n, 0.1), 2), ""), collapse = " ")
    if (line != PUBLISHED[[name]]) fail(paste0("Fader and Hardie (2010) Table 4 ", name, " not reproduced: ", line))
    out("metric.check.fader_hardie_2010_", name, "_derl_n5_to_n1=", line)
  }
}

true_surv <- sbg_survival(TRUE_ALPHA, TRUE_BETA, TERMS)
true_clv <- clv_new(true_surv)
true_rv <- residual_value(true_surv, CALIBRATION)

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

lifetimes <- draw_lifetimes(TRUE_ALPHA, TRUE_BETA, SUBSCRIBERS, HORIZON)
cal <- calibration_counts(lifetimes, CALIBRATION)
active <- c(SUBSCRIBERS, vapply(1:HORIZON, function(t) sum(lifetimes > t), integer(1)))
act <- function(t) active[t + 1]
out("metric.data.active_by_month=", paste0(0:HORIZON, ":", active, collapse = " "))
out("metric.data.retention_rate_calibration=",
    paste0(1:CALIBRATION, ":", vapply(1:CALIBRATION, function(t) fmt(act(t) / act(t - 1), 4), ""), collapse = " "))
out("metric.data.retention_rate_holdout=",
    paste0((CALIBRATION + 1):HORIZON, ":",
           vapply((CALIBRATION + 1):HORIZON, function(t) fmt(act(t) / act(t - 1), 4), ""), collapse = " "))

geo <- fit_geometric(cal$churned, cal$survivors, CALIBRATION)
sbg <- fit_sbg(cal$churned, cal$survivors, CALIBRATION)
se <- hessian_se(sbg$neg, sbg$x)
alpha <- sbg$alpha
beta <- sbg$beta
theta <- geo$theta
out("metric.fit.geometric theta=", fmt(theta, 4), " retention=", fmt(1 - theta, 4),
    " churners=", geo$deaths, " renewal_decisions=", geo$exposure, " loglik=", fmt(geo$ll, 3))
out("metric.fit.sbg alpha=", fmt(alpha, 4), " se=", fmt(alpha * se[1], 4), " beta=", fmt(beta, 4),
    " se=", fmt(beta * se[2], 4), " loglik=", fmt(sbg$ll, 3))
out("metric.fit.sbg mean_churn=", fmt(alpha / (alpha + beta), 4),
    " polarization=", fmt(1 / (alpha + beta + 1), 4))
out("metric.fit.likelihood_ratio=", fmt(2 * (sbg$ll - geo$ll), 3), " aic_geometric=", fmt(-2 * geo$ll + 2, 3),
    " aic_sbg=", fmt(-2 * sbg$ll + 4, 3))

last_theta <- 1 - act(CALIBRATION) / act(CALIBRATION - 1)
out("metric.fit.last_rate theta=", fmt(last_theta, 4), " retention=", fmt(1 - last_theta, 4))
s_sbg <- sbg_survival(alpha, beta, TERMS)
s_geo <- geometric_survival(theta, TERMS)
s_last <- vapply(0:TERMS, function(t) {
  if (t <= CALIBRATION) act(t) / SUBSCRIBERS
  else act(CALIBRATION) / SUBSCRIBERS * (1 - last_theta)^(t - CALIBRATION)
}, numeric(1))
for (t in c(6L, 12L, 18L, 24L)) {
  out("metric.projection.month=", t, " observed=", fmt(act(t) / SUBSCRIBERS, 4),
      " geometric=", fmt(s_geo[t + 1], 4), " last_rate=", fmt(s_last[t + 1], 4), " sbg=", fmt(s_sbg[t + 1], 4),
      " truth=", fmt(true_surv[t + 1], 4))
}
holdout <- (CALIBRATION + 1):HORIZON
observed <- vapply(holdout, function(t) act(t) / SUBSCRIBERS, numeric(1))
mae_geo <- plain_sum(abs(s_geo[holdout + 1] - observed)) / (HORIZON - CALIBRATION)
mae_sbg <- plain_sum(abs(s_sbg[holdout + 1] - observed)) / (HORIZON - CALIBRATION)
mae_last <- plain_sum(abs(s_last[holdout + 1] - observed)) / (HORIZON - CALIBRATION)
out("metric.projection.holdout_mean_abs_error_points geometric=", fmt(100 * mae_geo, 2),
    " last_rate=", fmt(100 * mae_last, 2), " sbg=", fmt(100 * mae_sbg, 2))
out("metric.projection.sbg_retention_rate=",
    paste0(c(1, 6, 12, 24), ":", vapply(c(1L, 6L, 12L, 24L), function(t) fmt(s_sbg[t + 1] / s_sbg[t], 4), ""),
           collapse = " "))

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)
survivors <- cal$survivors
out("metric.value.clv_new_subscriber geometric=", fmt(clv_geo, 2), " sbg=", fmt(clv_sbg, 2), " truth=", fmt(true_clv, 2))
out("metric.value.residual_per_survivor geometric=", fmt(rv_geo, 2), " last_rate=", fmt(rv_last, 2),
    " sbg=", fmt(rv_sbg, 2), " truth=", fmt(true_rv, 2))
out("metric.value.residual_cohort survivors=", survivors, " geometric=", fmt(survivors * rv_geo, 0),
    " last_rate=", fmt(survivors * rv_last, 0), " sbg=", fmt(survivors * rv_sbg, 0),
    " truth=", fmt(survivors * true_rv, 0))
out("metric.value.gap_vs_truth clv_new_geometric=", fmt(100 * (clv_geo / true_clv - 1), 1), "%",
    " clv_new_sbg=", fmt(100 * (clv_sbg / true_clv - 1), 1), "%",
    " residual_geometric=", fmt(100 * (rv_geo / true_rv - 1), 1), "%",
    " residual_last_rate=", fmt(100 * (rv_last / true_rv - 1), 1), "%",
    " residual_sbg=", fmt(100 * (rv_sbg / true_rv - 1), 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.
out("metric.horizon.share_of_value_beyond_month=", paste(vapply(c(24L, 60L, 120L), function(h) paste0(
  h, ":clv_sbg=", fmt(100 * (1 - clv_new(s_sbg, payments = h) / clv_sbg), 1), "%",
  ",residual_sbg=", fmt(100 * (1 - residual_value(s_sbg, CALIBRATION, last = h - 1L) / rv_sbg), 1), "%"), ""),
  collapse = " "))
out("metric.horizon.clv_first_payments=", paste(vapply(c(24L, 36L, 60L), function(p) paste0(
  p, ":geometric=", fmt(clv_new(s_geo, payments = p), 2), ",sbg=", fmt(clv_new(s_sbg, payments = p), 2),
  ",truth=", fmt(clv_new(true_surv, payments = p), 2)), ""), collapse = " "))

# The discount rate per period drives both the values and the size of the constant-rate gap.
out("metric.sensitivity.annual_equivalent_of_monthly_discount=", fmt(100 * ((1 + DISCOUNT)^12 - 1), 1), "%")
s_last_geo <- geometric_survival(last_theta, TERMS)
for (rate in c(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)
  out("metric.sensitivity.discount=", fmt(100 * rate, 1), "% clv_geometric=", fmt(clv_new(s_geo, rate), 2),
      " clv_sbg=", fmt(clv_new(s_sbg, rate), 2), " clv_truth=", fmt(t_clv, 2),
      " residual_sbg=", fmt(residual_value(s_sbg, CALIBRATION, rate), 2), " residual_truth=", fmt(t_rv, 2),
      " residual_gap_geometric=", fmt(100 * (g_rv / t_rv - 1), 1), "%",
      " residual_gap_last_rate=", fmt(100 * (l_rv / t_rv - 1), 1), "%")
}

boot_clv <- numeric(BOOTSTRAP)
boot_rv <- numeric(BOOTSTRAP)
for (b in seq_len(BOOTSTRAP)) {
  b_cal <- calibration_counts(draw_lifetimes(alpha, beta, SUBSCRIBERS, CALIBRATION), CALIBRATION)
  bf <- fit_sbg(b_cal$churned, b_cal$survivors, CALIBRATION)
  bs <- sbg_survival(bf$alpha, bf$beta, TERMS)
  boot_clv[b] <- clv_new(bs)
  boot_rv[b] <- residual_value(bs, CALIBRATION)
}
boot_clv <- sort(boot_clv)
boot_rv <- sort(boot_rv)
out("metric.uncertainty.parametric_bootstrap clv_new_95=", fmt(boot_clv[5], 2), "..", fmt(boot_clv[195], 2),
    " residual_per_survivor_95=", fmt(boot_rv[5], 2), "..", fmt(boot_rv[195], 2))

keys <- c("clv_geo", "clv_sbg", "rv_geo", "rv_last", "rv_sbg", "s24_geo", "s24_sbg")
rep <- matrix(0, REPLICATIONS, length(keys), dimnames = list(NULL, keys))
geo_below <- 0L
last_below <- 0L
for (r in seq_len(REPLICATIONS)) {
  r_cal <- calibration_counts(draw_lifetimes(TRUE_ALPHA, TRUE_BETA, SUBSCRIBERS, CALIBRATION), CALIBRATION)
  r_last <- r_cal$churned[CALIBRATION] / (r_cal$survivors + r_cal$churned[CALIBRATION])
  r_theta <- fit_geometric(r_cal$churned, r_cal$survivors, CALIBRATION)$theta
  rf <- fit_sbg(r_cal$churned, r_cal$survivors, CALIBRATION)
  rs <- sbg_survival(rf$alpha, rf$beta, TERMS)
  rg <- geometric_survival(r_theta, TERMS)
  rep[r, "clv_geo"] <- clv_new(rg)
  rep[r, "clv_sbg"] <- clv_new(rs)
  rep[r, "rv_geo"] <- residual_value(rg, CALIBRATION)
  rep[r, "rv_last"] <- residual_value(geometric_survival(r_last, TERMS), CALIBRATION)
  rep[r, "rv_sbg"] <- residual_value(rs, CALIBRATION)
  rep[r, "s24_geo"] <- rg[HORIZON + 1]
  rep[r, "s24_sbg"] <- rs[HORIZON + 1]
  if (rep[r, "rv_geo"] < true_rv) geo_below <- geo_below + 1L
  if (rep[r, "rv_last"] < true_rv) last_below <- last_below + 1L
}
truths <- c(true_clv, true_clv, true_rv, true_rv, true_rv, true_surv[HORIZON + 1], true_surv[HORIZON + 1])
digits <- c(2, 2, 2, 2, 2, 4, 4)
for (k in seq_along(keys)) {
  st <- summary_stats(rep[, k], truths[k])
  out("metric.monte_carlo.", keys[k], " mean=", fmt(st[1], digits[k]), " sd=", fmt(st[2], digits[k]),
      " bias=", fmt(st[3], digits[k]), " bias_mcse=", fmt(st[2] / sqrt(REPLICATIONS), digits[k]),
      " rmse=", fmt(st[4], digits[k]))
}
out("metric.monte_carlo.residual_below_truth geometric=", geo_below, " last_rate=", last_below, " of ", REPLICATIONS)
