#!/usr/bin/env Rscript
# MSC-P-039 choosing a statistical test. MIT License. Base R reference.
# The three procedures, the simulation scenarios, the nominal level and the
# tolerance are declared, and the datasets come from the declared recursion, so
# this file reproduces the same digits as the Python reference.
# Base R's pt() and pnorm() supply the tail probabilities; the Python reference
# computes them from the incomplete beta and the error function, and the two
# agree to well beyond the six printed decimals.

args <- commandArgs(trailingOnly = TRUE)
path <- if (length(args) >= 1) args[[1]] else file.path("..", "datasets", "msc-p039-two-group-outcome.csv")

raw_lines <- readLines(path, warn = FALSE)
raw_lines <- raw_lines[nzchar(raw_lines)]
if (raw_lines[[1]] != "unit_id,group,minutes") stop("exact schema required")
if (any(lengths(regmatches(raw_lines[-1], gregexpr(",", raw_lines[-1], fixed = TRUE))) != 2L)) stop("exactly three cells required per row")

d <- read.csv(path, stringsAsFactors = FALSE, colClasses = c("character", "character", "numeric"))
if (!identical(names(d), c("unit_id", "group", "minutes"))) stop("exact schema required")
if (nrow(d) < 20L) stop("too few sampling units for the declared comparison")
if (!identical(trimws(d$unit_id), sprintf("U%03d", seq_len(nrow(d))))) stop("units must be ordered U001, U002, ... without gaps")
if (!all(trimws(d$group) %in% c("control", "treated"))) stop("the group column must contain only control and treated")
if (any(!is.finite(d$minutes)) || any(d$minutes <= 0)) stop("the outcome must be finite and strictly positive")
control <- d$minutes[trimws(d$group) == "control"]
treated <- d$minutes[trimws(d$group) == "treated"]
if (length(control) < 10L || length(treated) < 10L) stop("each declared group needs at least ten sampling units")

NOMINAL <- 0.05
REPLICATIONS <- 2000L
PER_GROUP <- 25L
SKEW_LOG_SD <- 0.85
ALTERNATIVE_LOG_SHIFT <- 0.45
NORMALITY_LEVEL <- 0.05
LEVEL_TOLERANCE <- 0.0125
LCG_SEED <- 20260915
LCG_MULTIPLIER <- 1103515245
LCG_INCREMENT <- 12345
LCG_MODULUS <- 2147483648
# R numbers are doubles, exact for integers only up to 2^53. Written directly,
# LCG_MULTIPLIER * state reaches 2.4e18 and would be rounded. Splitting the
# multiplier keeps every product exact: 16838 * 65536 + 20077 is the multiplier.
LCG_HIGH <- 16838
LCG_LOW <- 20077

welch <- function(first, second) {
  n1 <- length(first); n2 <- length(second)
  m1 <- mean(first); m2 <- mean(second)
  v1 <- var(first); v2 <- var(second)
  if (v1 <= 0 && v2 <= 0) stop("both declared groups are constant")
  statistic <- (m2 - m1) / sqrt(v1 / n1 + v2 / n2)
  degrees <- (v1 / n1 + v2 / n2)^2 / ((v1 / n1)^2 / (n1 - 1) + (v2 / n2)^2 / (n2 - 1))
  c(statistic, degrees, 2 * pt(-abs(statistic), degrees))
}

mann_whitney <- function(first, second) {
  combined <- c(first, second)
  labels <- c(rep(0L, length(first)), rep(1L, length(second)))
  ranks <- rank(combined, ties.method = "average")
  n1 <- length(first); n2 <- length(second); total <- n1 + n2
  statistic <- sum(ranks[labels == 0L]) - n1 * (n1 + 1) / 2
  counts <- table(combined)
  correction <- sum(counts^3 - counts)
  variance <- n1 * n2 / 12 * ((total + 1) - correction / (total * (total - 1)))
  if (variance <= 0) stop("the declared rank test has no variance")
  standardized <- (statistic - n1 * n2 / 2) / sqrt(variance)
  c(statistic, standardized, 2 * pnorm(-abs(standardized)))
}

jarque_bera <- function(values) {
  n <- length(values)
  centred <- values - mean(values)
  second <- mean(centred^2)
  if (second <= 0) stop("a declared group is constant, so normality cannot be assessed")
  skewness <- mean(centred^3) / second^1.5
  kurtosis <- mean(centred^4) / second^2
  statistic <- n / 6 * (skewness^2 + (kurtosis - 3)^2 / 4)
  c(statistic, exp(-statistic / 2))
}

pretest <- function(first, second) {
  left <- jarque_bera(first)[[2]]
  right <- jarque_bera(second)[[2]]
  if (left > NORMALITY_LEVEL && right > NORMALITY_LEVEL) {
    list(choice = "welch", probability = welch(first, second)[[3]])
  } else {
    list(choice = "mann-whitney", probability = mann_whitney(first, second)[[3]])
  }
}

lcg_state <- LCG_SEED
lcg_spare <- NULL
lcg_reset <- function() { lcg_state <<- LCG_SEED; lcg_spare <<- NULL }
lcg_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
}
lcg_normal <- function() {
  if (!is.null(lcg_spare)) {
    value <- lcg_spare
    lcg_spare <<- NULL
    return(value)
  }
  first <- lcg_uniform(); second <- lcg_uniform()
  radius <- sqrt(-2 * log(first)); angle <- 2 * pi * second
  lcg_spare <<- radius * sin(angle)
  radius * cos(angle)
}

simulate <- function(distribution, shift) {
  lcg_reset()
  counts <- c(welch = 0L, rank = 0L, pretest = 0L)
  chose_welch <- 0L
  for (replication in seq_len(REPLICATIONS)) {
    first <- numeric(PER_GROUP); second <- numeric(PER_GROUP)
    for (i in seq_len(PER_GROUP)) {
      if (distribution == "normal") {
        first[[i]] <- lcg_normal()
        second[[i]] <- lcg_normal() + shift
      } else {
        first[[i]] <- exp(SKEW_LOG_SD * lcg_normal())
        second[[i]] <- exp(SKEW_LOG_SD * lcg_normal() + shift)
      }
    }
    if (welch(first, second)[[3]] < NOMINAL) counts[["welch"]] <- counts[["welch"]] + 1L
    if (mann_whitney(first, second)[[3]] < NOMINAL) counts[["rank"]] <- counts[["rank"]] + 1L
    decision <- pretest(first, second)
    if (decision$choice == "welch") chose_welch <- chose_welch + 1L
    if (decision$probability < NOMINAL) counts[["pretest"]] <- counts[["pretest"]] + 1L
  }
  list(rates = counts / REPLICATIONS, chose_welch = chose_welch / REPLICATIONS)
}

w <- welch(control, treated)
u <- mann_whitney(control, treated)
control_jb <- jarque_bera(control)
treated_jb <- jarque_bera(treated)
decision <- pretest(control, treated)

cat(sprintf("control_units=%d\n", length(control)))
cat(sprintf("treated_units=%d\n", length(treated)))
cat(sprintf("welch_statistic=%.6f\n", w[[1]]))
cat(sprintf("welch_degrees_of_freedom=%.6f\n", w[[2]]))
cat(sprintf("welch_probability=%.6f\n", w[[3]]))
cat(sprintf("rank_statistic=%.6f\n", u[[1]]))
cat(sprintf("rank_standardized=%.6f\n", u[[2]]))
cat(sprintf("rank_probability=%.6f\n", u[[3]]))
cat(sprintf("control_normality_statistic=%.6f\n", control_jb[[1]]))
cat(sprintf("control_normality_probability=%.6f\n", control_jb[[2]]))
cat(sprintf("treated_normality_statistic=%.6f\n", treated_jb[[1]]))
cat(sprintf("treated_normality_probability=%.6f\n", treated_jb[[2]]))
cat(sprintf("pretest_choice=%s\n", decision$choice))
cat(sprintf("pretest_probability=%.6f\n", decision$probability))

cat(sprintf("replications=%d\n", REPLICATIONS))
cat(sprintf("units_per_group=%d\n", PER_GROUP))
cat(sprintf("nominal_level=%.2f\n", NOMINAL))
scenarios <- list(
  list(label = "normal_null", distribution = "normal", shift = 0),
  list(label = "skewed_null", distribution = "skewed", shift = 0),
  list(label = "skewed_alternative", distribution = "skewed", shift = ALTERNATIVE_LOG_SHIFT)
)
levels_by_label <- list()
for (scenario in scenarios) {
  outcome <- simulate(scenario$distribution, scenario$shift)
  cat(sprintf("%s_welch=%.6f\n", scenario$label, outcome$rates[["welch"]]))
  cat(sprintf("%s_mann_whitney=%.6f\n", scenario$label, outcome$rates[["rank"]]))
  cat(sprintf("%s_pretest=%.6f\n", scenario$label, outcome$rates[["pretest"]]))
  cat(sprintf("%s_pretest_chose_welch=%.6f\n", scenario$label, outcome$chose_welch))
  if (scenario$shift == 0) levels_by_label[[scenario$label]] <- outcome$rates
}

pretest_honest <- TRUE
for (label in c("normal_null", "skewed_null")) {
  for (name in c("welch", "rank", "pretest")) {
    honest <- abs(levels_by_label[[label]][[name]] - NOMINAL) <= LEVEL_TOLERANCE
    printed <- if (name == "rank") "mann_whitney" else name
    cat(sprintf("%s_%s_flag=%s\n", label, printed, if (honest) "PASS" else "FAIL"))
    if (name == "pretest" && !honest) pretest_honest <- FALSE
  }
}
cat(sprintf("level_tolerance=%.4f\n", LEVEL_TOLERANCE))
cat(sprintf("verdict=%s\n", if (pretest_honest) "PRETEST_RULE_HOLDS_ITS_LEVEL" else "PRETEST_RULE_DOES_NOT_HOLD_ITS_LEVEL"))
