#!/usr/bin/env Rscript
# MSC-P-035 saturation and adstock. MIT License. Base R reference implementation.
# The grid, the region tolerance and both thresholds are declared, so this file
# reproduces the same digits as the Python reference. Nothing is random.

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

raw_lines <- readLines(path, warn = FALSE)
raw_lines <- raw_lines[nzchar(raw_lines)]
if (raw_lines[[1]] != "week,media_spend_keur,revenue_keur") 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("integer", "numeric", "numeric"))
if (!identical(names(d), c("week", "media_spend_keur", "revenue_keur"))) stop("exact schema required")
if (nrow(d) < 104L) stop("too few weeks for the declared grid search")
if (!identical(d$week, seq_len(nrow(d)))) stop("weeks must be numbered 1, 2, ... without gaps")
if (any(!is.finite(d$media_spend_keur)) || any(d$media_spend_keur <= 0)) stop("media spend must be finite and strictly positive")
if (any(!is.finite(d$revenue_keur)) || any(d$revenue_keur <= 0)) stop("revenue must be finite and strictly positive")

PERIOD <- 52L
HARMONICS <- 2L
CARRYOVER_GRID <- round(seq(0, 0.9, by = 0.1), 1)
SHAPE_GRID <- c(0.6, 1.0, 1.4, 1.8, 2.2, 2.6, 3.0)
HALF_SATURATION_GRID <- c(25, 35, 45, 55, 65, 75, 85, 95)
TRUE_PARAMETERS <- c(0.6, 1.8, 55)
REGION_TOLERANCE <- 0.01
CARRYOVER_RANGE_MAX <- 0.20
MARGINAL_RATIO_MAX <- 1.50

x <- d$media_spend_keur
y <- d$revenue_keur
n <- length(y)
weeks <- seq_len(n)

columns <- list(rep(1, n), as.numeric(weeks))
for (k in seq_len(HARMONICS)) {
  angle <- 2 * pi * k * weeks / PERIOD
  columns <- c(columns, list(sin(angle)), list(cos(angle)))
}
BASE <- do.call(cbind, columns)

adstock <- function(carryover) {
  out <- numeric(n)
  carried <- 0
  for (i in seq_len(n)) {
    carried <- x[[i]] + carryover * carried
    out[[i]] <- carried
  }
  out
}

saturate <- function(series, shape, half_saturation) {
  denominator <- half_saturation^shape
  series^shape / (denominator + series^shape)
}

marginal_slope <- function(level, shape, half_saturation) {
  denominator <- half_saturation^shape
  shape * denominator * level^(shape - 1) / (denominator + level^shape)^2
}

carried_by_rate <- lapply(CARRYOVER_GRID, adstock)
names(carried_by_rate) <- as.character(CARRYOVER_GRID)

grid <- list()
for (carryover in CARRYOVER_GRID) {
  series <- carried_by_rate[[as.character(carryover)]]
  mean_level <- mean(series)
  for (shape in SHAPE_GRID) {
    for (half_saturation in HALF_SATURATION_GRID) {
      X <- cbind(BASE, saturate(series, shape, half_saturation))
      normal <- crossprod(X)
      if (rcond(normal) < 1e-12) stop("the declared design matrix is singular for this grid point")
      coefficients <- as.vector(solve(normal, crossprod(X, y)))
      residual <- sum((y - as.vector(X %*% coefficients))^2)
      amplitude <- coefficients[[length(coefficients)]]
      grid[[length(grid) + 1L]] <- list(
        carryover = carryover, shape = shape, half_saturation = half_saturation,
        residual = residual, amplitude = amplitude,
        marginal = amplitude * marginal_slope(mean_level, shape, half_saturation)
      )
    }
  }
}

residuals_all <- vapply(grid, function(e) e$residual, numeric(1))
order_index <- order(residuals_all,
                     vapply(grid, function(e) e$carryover, numeric(1)),
                     vapply(grid, function(e) e$shape, numeric(1)),
                     vapply(grid, function(e) e$half_saturation, numeric(1)))
best <- grid[[order_index[[1]]]]
region <- Filter(function(e) e$residual <= best$residual * (1 + REGION_TOLERANCE), grid)

carryovers <- vapply(region, function(e) e$carryover, numeric(1))
shapes <- vapply(region, function(e) e$shape, numeric(1))
saturations <- vapply(region, function(e) e$half_saturation, numeric(1))
marginals <- vapply(region, function(e) e$marginal, numeric(1))
if (min(marginals) <= 0) stop("a grid point in the region implies a non-positive marginal return")

carryover_range <- max(carryovers) - min(carryovers)
marginal_ratio <- max(marginals) / min(marginals)
carryover_flag <- if (carryover_range <= CARRYOVER_RANGE_MAX) "PASS" else "FAIL"
marginal_flag <- if (marginal_ratio <= MARGINAL_RATIO_MAX) "PASS" else "FAIL"
verdict <- if (carryover_flag == "PASS" && marginal_flag == "PASS") "RESPONSE_CURVE_IDENTIFIED_FOR_REALLOCATION" else "DIAGNOSTIC_BLOCKS_RESPONSE_CURVE"

matches_truth <- function(e) {
  isTRUE(all.equal(c(e$carryover, e$shape, e$half_saturation), TRUE_PARAMETERS))
}
truth <- Filter(matches_truth, grid)[[1]]
truth_inside <- any(vapply(region, matches_truth, logical(1)))

total <- sum((y - mean(y))^2)
cat(sprintf("weeks=%d\n", n))
cat(sprintf("grid_points=%d\n", length(grid)))
cat(sprintf("best_carryover=%.6f\n", best$carryover))
cat(sprintf("best_shape=%.6f\n", best$shape))
cat(sprintf("best_half_saturation=%.6f\n", best$half_saturation))
cat(sprintf("best_amplitude=%.6f\n", best$amplitude))
cat(sprintf("r_squared=%.6f\n", 1 - best$residual / total))
cat(sprintf("residual_sd=%.6f\n", sqrt(best$residual / (n - ncol(BASE) - 1))))
cat(sprintf("region_tolerance=%.2f\n", REGION_TOLERANCE))
cat(sprintf("region_points=%d\n", length(region)))
cat(sprintf("carryover_min=%.6f\n", min(carryovers)))
cat(sprintf("carryover_max=%.6f\n", max(carryovers)))
cat(sprintf("carryover_range=%.6f\n", carryover_range))
cat(sprintf("carryover_flag=%s\n", carryover_flag))
cat(sprintf("shape_min=%.6f\n", min(shapes)))
cat(sprintf("shape_max=%.6f\n", max(shapes)))
cat(sprintf("half_saturation_min=%.6f\n", min(saturations)))
cat(sprintf("half_saturation_max=%.6f\n", max(saturations)))
cat(sprintf("marginal_min=%.6f\n", min(marginals)))
cat(sprintf("marginal_max=%.6f\n", max(marginals)))
cat(sprintf("marginal_ratio=%.6f\n", marginal_ratio))
cat(sprintf("marginal_flag=%s\n", marginal_flag))
cat(sprintf("best_marginal_return=%.6f\n", best$marginal))
cat(sprintf("true_parameters_inside_region=%s\n", if (truth_inside) "YES" else "NO"))
cat(sprintf("true_parameters_fit_excess_percent=%.6f\n", 100 * (truth$residual / best$residual - 1)))
cat(sprintf("true_marginal_return=%.6f\n", truth$marginal))
cat(sprintf("verdict=%s\n", verdict))
