#!/usr/bin/env Rscript
# MSC-P-032 segmentation stability. MIT License. Base R reference implementation.
# Resamples come from the declared linear congruential recursion, never from the
# R generator, so this file reproduces the same digits as the Python reference.

args <- commandArgs(trailingOnly = TRUE)
path <- if (length(args) >= 1) args[[1]] else file.path("..", "datasets", "msc-p032-stability-panel.csv")

raw_lines <- readLines(path, warn = FALSE)
raw_lines <- raw_lines[nzchar(raw_lines)]
if (raw_lines[[1]] != "customer_id,recency_days,frequency_12m,avg_basket_eur") stop("exact schema required")
if (any(lengths(regmatches(raw_lines[-1], gregexpr(",", raw_lines[-1], fixed = TRUE))) != 3L)) stop("exactly four cells required per row")

d <- read.csv(path, stringsAsFactors = FALSE, colClasses = c("character", "numeric", "numeric", "numeric"))
required <- c("customer_id", "recency_days", "frequency_12m", "avg_basket_eur")
if (!identical(names(d), required)) stop("exact schema required")
if (nrow(d) < 100L) stop("too few customers for the declared stability analysis")
if (!identical(trimws(d$customer_id), sprintf("T%04d", seq_len(nrow(d))))) stop("customers must be ordered T0001, T0002, ... without gaps")
positive <- cbind(d$recency_days, d$frequency_12m, d$avg_basket_eur)
if (any(!is.finite(positive)) || any(positive <= 0)) stop("recency, frequency and basket must be finite and strictly positive")

CANDIDATES <- c(2L, 3L, 4L)
PROPOSED <- 4L
RESAMPLES <- 40L
MAX_ITERATIONS <- 100L
RECOVERY_MIN <- 0.75
MARGIN_MIN <- 0.10
LCG_SEED <- 20260912
LCG_MULTIPLIER <- 1103515245
LCG_INCREMENT <- 12345
LCG_MODULUS <- 2147483648
# R numbers are doubles, which hold integers exactly only up to 2^53. Written
# directly, LCG_MULTIPLIER * state reaches 2.4e18 and would be rounded, giving a
# different recursion from the Python reference. Splitting the multiplier keeps
# every intermediate product below 2^53 and is exact: 16838 * 65536 + 20077 is
# LCG_MULTIPLIER itself.
LCG_HIGH <- 16838
LCG_LOW <- 20077

X <- cbind(log(d$recency_days), d$frequency_12m, log(d$avg_basket_eur))
centres <- colMeans(X)
spreads <- apply(X, 2, sd)
if (any(spreads <= 0)) stop("a declared feature has no variation")
Z <- sweep(sweep(X, 2, centres, "-"), 2, spreads, "/")
n <- nrow(Z)

lcg_state <- LCG_SEED
lcg_reset <- function() lcg_state <<- LCG_SEED
lcg_next <- function(size) {
  high <- (LCG_HIGH * lcg_state) %% LCG_MODULUS
  lcg_state <<- (high * 65536 + LCG_LOW * lcg_state + LCG_INCREMENT) %% LCG_MODULUS
  as.integer(lcg_state %% size)
}

start_centroids <- function(points, k) {
  order_index <- order(rowSums(points), seq_len(nrow(points)))
  positions <- sapply(seq_len(k), function(j) round((j - 0.5) / k * (nrow(points) - 1))) + 1L
  points[order_index[positions], , drop = FALSE]
}

assign_labels <- function(points, centroids) {
  apply(points, 1, function(row) which.min(rowSums((centroids - matrix(row, nrow = nrow(centroids), ncol = length(row), byrow = TRUE))^2)))
}

cluster <- function(points, k) {
  centroids <- start_centroids(points, k)
  labels <- rep(0L, nrow(points))
  for (iteration in seq_len(MAX_ITERATIONS)) {
    new_labels <- assign_labels(points, centroids)
    changed <- !identical(new_labels, labels)
    labels <- new_labels
    for (j in seq_len(k)) {
      members <- points[labels == j, , drop = FALSE]
      if (nrow(members) == 0L) stop("the declared segmentation collapsed to fewer segments")
      centroids[j, ] <- colMeans(members)
    }
    if (!changed) break
  }
  list(labels = labels, centroids = centroids)
}

jaccard <- function(a, b) {
  union_size <- sum(a | b)
  if (union_size == 0) 0 else sum(a & b) / union_size
}

recovery <- function(points, k) {
  solution <- cluster(points, k)
  original <- lapply(seq_len(k), function(j) solution$labels == j)
  totals <- numeric(k)
  lcg_reset()
  for (b in seq_len(RESAMPLES)) {
    index <- vapply(seq_len(nrow(points)), function(i) lcg_next(nrow(points)) + 1L, integer(1))
    rebuilt <- cluster(points[index, , drop = FALSE], k)
    labels <- assign_labels(points, rebuilt$centroids)
    for (j in seq_len(k)) {
      totals[j] <- totals[j] + max(vapply(seq_len(k), function(m) jaccard(original[[j]], labels == m), numeric(1)))
    }
  }
  totals / RESAMPLES
}

scrambled <- function(points) {
  lcg_reset()
  out <- points
  for (column in seq_len(ncol(points))) {
    values <- points[, column]
    for (i in seq(length(values), 2L)) {
      swap <- lcg_next(i) + 1L
      temporary <- values[[i]]
      values[[i]] <- values[[swap]]
      values[[swap]] <- temporary
    }
    out[, column] <- values
  }
  out
}

reference <- scrambled(Z)
cat(sprintf("customers=%d\n", n))
cat(sprintf("resamples=%d\n", RESAMPLES))
cat(sprintf("proposed_segments=%d\n", PROPOSED))
flags <- list()
for (k in CANDIDATES) {
  observed <- recovery(Z, k)
  null_recovery <- recovery(reference, k)
  weakest <- min(observed)
  null_weakest <- min(null_recovery)
  margin <- weakest - null_weakest
  flag <- if (weakest >= RECOVERY_MIN && margin >= MARGIN_MIN) "STABLE" else "NOT_STABLE"
  flags[[as.character(k)]] <- flag
  cat(sprintf("recovery_k%d=%s\n", k, paste(sprintf("%.6f", observed), collapse = ",")))
  cat(sprintf("weakest_k%d=%.6f\n", k, weakest))
  cat(sprintf("null_weakest_k%d=%.6f\n", k, null_weakest))
  cat(sprintf("margin_k%d=%.6f\n", k, margin))
  cat(sprintf("flag_k%d=%s\n", k, flag))
}
verdict <- if (identical(flags[[as.character(PROPOSED)]], "STABLE")) "PROPOSED_SEGMENTATION_STABLE" else "PROPOSED_SEGMENTATION_NOT_STABLE"
cat(sprintf("verdict=%s\n", verdict))
