#!/usr/bin/env python3
"""MSC-P-035 saturation and adstock. MIT License.

Fits a declared media response model — geometric carryover followed by a Hill
saturation — on a declared grid, then asks the question that decides whether the
fit may be used to move budget: how many other parameter combinations explain
the data almost as well, and do they agree on the return of one more unit of
spend?

The dossier separates two things that are routinely confused. A model can track
the outcome closely and still leave its own parameters undetermined, because
carryover, curvature and half-saturation trade off against one another. Fit is
reported first; identification is reported second; only the second one licenses
a reallocation.

Nothing is random here: the grid, the tolerance and both thresholds are
declared, so every implementation reproduces the same digits.
"""
import csv
import math
import sys
from pathlib import Path

REQUIRED = ["week", "media_spend_keur", "revenue_keur"]
PERIOD = 52
HARMONICS = 2
CARRYOVER_GRID = [round(0.1 * step, 1) for step in range(0, 10)]
SHAPE_GRID = [0.6, 1.0, 1.4, 1.8, 2.2, 2.6, 3.0]
HALF_SATURATION_GRID = [25.0, 35.0, 45.0, 55.0, 65.0, 75.0, 85.0, 95.0]
TRUE_PARAMETERS = (0.6, 1.8, 55.0)
REGION_TOLERANCE = 0.01
CARRYOVER_RANGE_MAX = 0.20
MARGINAL_RATIO_MAX = 1.50


def load(path: Path):
    with path.open(newline="", encoding="utf-8") as handle:
        reader = csv.DictReader(handle)
        if reader.fieldnames != REQUIRED:
            raise ValueError("exact schema required")
        rows = list(reader)
    if len(rows) < 104:
        raise ValueError("too few weeks for the declared grid search")
    spend, revenue = [], []
    for index, row in enumerate(rows, start=1):
        if set(row) != set(REQUIRED) or any(row[key] is None for key in REQUIRED):
            raise ValueError("exactly three cells required per row")
        if row["week"].strip() != str(index):
            raise ValueError("weeks must be numbered 1, 2, ... without gaps")
        x = float(row["media_spend_keur"])
        y = float(row["revenue_keur"])
        if not math.isfinite(x) or x <= 0:
            raise ValueError("media spend must be finite and strictly positive")
        if not math.isfinite(y) or y <= 0:
            raise ValueError("revenue must be finite and strictly positive")
        spend.append(x)
        revenue.append(y)
    return spend, revenue


def adstock(spend, carryover):
    """Declared carryover: A(t) = x(t) + carryover * A(t-1), starting from A(0) = 0."""
    carried = 0.0
    series = []
    for value in spend:
        carried = value + carryover * carried
        series.append(carried)
    return series


def saturate(series, shape, half_saturation):
    """Declared Hill saturation, reported as a share between zero and one."""
    denominator = half_saturation ** shape
    return [value ** shape / (denominator + value ** shape) for value in series]


def marginal(level, shape, half_saturation):
    """Derivative of the declared saturation with respect to the adstock level."""
    denominator = half_saturation ** shape
    return shape * denominator * level ** (shape - 1.0) / (denominator + level ** shape) ** 2


def solve(matrix, vector):
    """Gaussian elimination with partial pivoting on the normal equations."""
    size = len(vector)
    augmented = [list(row) + [value] for row, value in zip(matrix, vector)]
    for column in range(size):
        pivot = max(range(column, size), key=lambda index: abs(augmented[index][column]))
        if abs(augmented[pivot][column]) < 1e-12:
            raise ValueError("the declared design matrix is singular for this grid point")
        augmented[column], augmented[pivot] = augmented[pivot], augmented[column]
        for index in range(column + 1, size):
            factor = augmented[index][column] / augmented[column][column]
            for position in range(column, size + 1):
                augmented[index][position] -= factor * augmented[column][position]
    solution = [0.0] * size
    for column in range(size - 1, -1, -1):
        total = augmented[column][size] - sum(augmented[column][position] * solution[position] for position in range(column + 1, size))
        solution[column] = total / augmented[column][column]
    return solution


def baseline_columns(count):
    """Declared baseline: intercept, linear trend and two annual harmonics."""
    columns = [[1.0] * count, [float(week) for week in range(1, count + 1)]]
    for index in range(1, HARMONICS + 1):
        angles = [2.0 * math.pi * index * week / PERIOD for week in range(1, count + 1)]
        columns.append([math.sin(angle) for angle in angles])
        columns.append([math.cos(angle) for angle in angles])
    return columns


def fit(columns, revenue):
    rows = list(zip(*columns))
    size = len(columns)
    normal = [[sum(row[i] * row[j] for row in rows) for j in range(size)] for i in range(size)]
    right = [sum(row[i] * value for row, value in zip(rows, revenue)) for i in range(size)]
    coefficients = solve(normal, right)
    residual = sum((value - sum(a * b for a, b in zip(row, coefficients))) ** 2 for row, value in zip(rows, revenue))
    return coefficients, residual


def analyze(spend, revenue):
    count = len(revenue)
    base = baseline_columns(count)
    mean_revenue = sum(revenue) / count
    total = sum((value - mean_revenue) ** 2 for value in revenue)
    carried = {carryover: adstock(spend, carryover) for carryover in CARRYOVER_GRID}
    grid = []
    for carryover in CARRYOVER_GRID:
        series = carried[carryover]
        mean_level = sum(series) / count
        for shape in SHAPE_GRID:
            for half_saturation in HALF_SATURATION_GRID:
                saturated = saturate(series, shape, half_saturation)
                coefficients, residual = fit(base + [saturated], revenue)
                grid.append({
                    "carryover": carryover, "shape": shape, "half_saturation": half_saturation,
                    "residual": residual, "amplitude": coefficients[-1],
                    "marginal": coefficients[-1] * marginal(mean_level, shape, half_saturation),
                })
    best = min(grid, key=lambda entry: (entry["residual"], entry["carryover"], entry["shape"], entry["half_saturation"]))
    region = [entry for entry in grid if entry["residual"] <= best["residual"] * (1.0 + REGION_TOLERANCE)]
    carryovers = [entry["carryover"] for entry in region]
    shapes = [entry["shape"] for entry in region]
    saturations = [entry["half_saturation"] for entry in region]
    marginals = [entry["marginal"] for entry in region]
    if min(marginals) <= 0:
        raise ValueError("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 = "PASS" if carryover_range <= CARRYOVER_RANGE_MAX else "FAIL"
    marginal_flag = "PASS" if marginal_ratio <= MARGINAL_RATIO_MAX else "FAIL"
    verdict = "RESPONSE_CURVE_IDENTIFIED_FOR_REALLOCATION" if carryover_flag == "PASS" and marginal_flag == "PASS" else "DIAGNOSTIC_BLOCKS_RESPONSE_CURVE"
    truth = next(
        entry for entry in grid
        if (entry["carryover"], entry["shape"], entry["half_saturation"]) == TRUE_PARAMETERS
    )
    truth_inside = truth in region
    return {
        "truth_excess": 100.0 * (truth["residual"] / best["residual"] - 1.0),
        "truth_marginal": truth["marginal"],
        "count": count, "grid": len(grid), "best": best,
        "r_squared": 1.0 - best["residual"] / total,
        "residual_sd": math.sqrt(best["residual"] / (count - len(base) - 1)),
        "region": len(region),
        "carryover_min": min(carryovers), "carryover_max": max(carryovers), "carryover_range": carryover_range,
        "shape_min": min(shapes), "shape_max": max(shapes),
        "half_saturation_min": min(saturations), "half_saturation_max": max(saturations),
        "marginal_min": min(marginals), "marginal_max": max(marginals), "marginal_ratio": marginal_ratio,
        "carryover_flag": carryover_flag, "marginal_flag": marginal_flag,
        "truth_inside": truth_inside, "verdict": verdict,
    }


def main():
    path = Path(sys.argv[1]) if len(sys.argv) > 1 else Path(__file__).resolve().parent.parent / "datasets" / "msc-p035-media-series.csv"
    spend, revenue = load(path)
    result = analyze(spend, revenue)
    best = result["best"]
    print(f"weeks={result['count']}")
    print(f"grid_points={result['grid']}")
    print(f"best_carryover={best['carryover']:.6f}")
    print(f"best_shape={best['shape']:.6f}")
    print(f"best_half_saturation={best['half_saturation']:.6f}")
    print(f"best_amplitude={best['amplitude']:.6f}")
    print(f"r_squared={result['r_squared']:.6f}")
    print(f"residual_sd={result['residual_sd']:.6f}")
    print(f"region_tolerance={REGION_TOLERANCE:.2f}")
    print(f"region_points={result['region']}")
    print(f"carryover_min={result['carryover_min']:.6f}")
    print(f"carryover_max={result['carryover_max']:.6f}")
    print(f"carryover_range={result['carryover_range']:.6f}")
    print(f"carryover_flag={result['carryover_flag']}")
    print(f"shape_min={result['shape_min']:.6f}")
    print(f"shape_max={result['shape_max']:.6f}")
    print(f"half_saturation_min={result['half_saturation_min']:.6f}")
    print(f"half_saturation_max={result['half_saturation_max']:.6f}")
    print(f"marginal_min={result['marginal_min']:.6f}")
    print(f"marginal_max={result['marginal_max']:.6f}")
    print(f"marginal_ratio={result['marginal_ratio']:.6f}")
    print(f"marginal_flag={result['marginal_flag']}")
    print(f"best_marginal_return={best['marginal']:.6f}")
    print(f"true_parameters_inside_region={'YES' if result['truth_inside'] else 'NO'}")
    print(f"true_parameters_fit_excess_percent={result['truth_excess']:.6f}")
    print(f"true_marginal_return={result['truth_marginal']:.6f}")
    print(f"verdict={result['verdict']}")


if __name__ == "__main__":
    main()
