#!/usr/bin/env python3
"""MSC-P-033 marketing forecast validation. MIT License.

Validates a declared forecasting model the way it would be used: refitted at
every forecast origin on the data available up to that origin, then judged on
weeks it has never seen. Two questions are asked separately and answered
separately. Does the model beat a seasonal naive benchmark on the same scale,
which is what the mean absolute scaled error measures? And do its prediction
intervals contain the outcome as often as they claim?

Nothing is random here: the design matrix, the origins, the horizon, the
nominal level and both thresholds are declared, so every implementation
reproduces the same digits.
"""
import csv
import math
import sys
from pathlib import Path

REQUIRED = ["week", "revenue_eur"]
PERIOD = 52
HARMONICS = 2
FIRST_ORIGIN = 104
ORIGIN_STEP = 4
LAST_ORIGIN = 152
HORIZON = 4
NOMINAL = 0.80
NORMAL_QUANTILE = 1.281552  # the 0.90 quantile of the standard normal law
MASE_MAX = 1.00
COVERAGE_TOLERANCE = 0.10


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) < LAST_ORIGIN + HORIZON:
        raise ValueError("the series is shorter than the declared validation design")
    series = []
    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 two cells required per row")
        if row["week"].strip() != str(index):
            raise ValueError("weeks must be numbered 1, 2, ... without gaps")
        value = float(row["revenue_eur"])
        if not math.isfinite(value) or value <= 0:
            raise ValueError("the outcome must be finite and strictly positive")
        series.append(value)
    return series


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


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 on this window")
        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 fit(series, origin):
    """Least squares on weeks 1..origin only; nothing beyond the origin is read."""
    rows = [design_row(week) for week in range(1, origin + 1)]
    size = len(rows[0])
    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, series[:origin])) for i in range(size)]
    coefficients = solve(normal, right)
    residuals = [value - sum(a * b for a, b in zip(row, coefficients)) for row, value in zip(rows, series[:origin])]
    degrees = origin - size
    if degrees <= 0:
        raise ValueError("too few observations before the origin for the declared design")
    deviation = math.sqrt(sum(residual ** 2 for residual in residuals) / degrees)
    return coefficients, deviation


def scaling(series, origin):
    """Declared scale: in-sample mean absolute seasonal naive error up to the origin."""
    errors = [abs(series[index] - series[index - PERIOD]) for index in range(PERIOD, origin)]
    if not errors:
        raise ValueError("the training window is too short for the declared seasonal scale")
    scale = sum(errors) / len(errors)
    if scale <= 0:
        raise ValueError("the declared seasonal scale is zero")
    return scale


def interval_score(lower, upper, actual, alpha):
    """Winkler interval score: width, plus a penalty proportional to any miss."""
    score = upper - lower
    if actual < lower:
        score += 2.0 / alpha * (lower - actual)
    elif actual > upper:
        score += 2.0 / alpha * (actual - upper)
    return score


def analyze(series):
    alpha = 1.0 - NOMINAL
    origins = list(range(FIRST_ORIGIN, LAST_ORIGIN + 1, ORIGIN_STEP))
    scaled_model = []
    scaled_benchmark = []
    by_horizon = {horizon: [] for horizon in range(1, HORIZON + 1)}
    covered = 0
    widths = []
    scores = []
    absolute = []
    for origin in origins:
        coefficients, deviation = fit(series, origin)
        scale = scaling(series, origin)
        half_width = NORMAL_QUANTILE * deviation
        for horizon in range(1, HORIZON + 1):
            week = origin + horizon
            actual = series[week - 1]
            prediction = sum(a * b for a, b in zip(design_row(week), coefficients))
            error = abs(actual - prediction)
            absolute.append(error)
            scaled_model.append(error / scale)
            by_horizon[horizon].append(error / scale)
            scaled_benchmark.append(abs(actual - series[week - 1 - PERIOD]) / scale)
            lower, upper = prediction - half_width, prediction + half_width
            if lower <= actual <= upper:
                covered += 1
            widths.append(upper - lower)
            scores.append(interval_score(lower, upper, actual, alpha))
    count = len(scaled_model)
    mase = sum(scaled_model) / count
    benchmark = sum(scaled_benchmark) / count
    coverage = covered / count
    accuracy_flag = "PASS" if mase < MASE_MAX else "FAIL"
    calibration_flag = "PASS" if abs(coverage - NOMINAL) <= COVERAGE_TOLERANCE else "FAIL"
    verdict = "FORECAST_READABLE_FOR_PLANNING" if accuracy_flag == "PASS" and calibration_flag == "PASS" else "DIAGNOSTIC_BLOCKS_FORECAST_READING"
    return {
        "origins": origins, "forecasts": count, "mase": mase, "benchmark": benchmark,
        "mae": sum(absolute) / count, "coverage": coverage,
        "width": sum(widths) / count, "score": sum(scores) / count,
        "by_horizon": {horizon: sum(values) / len(values) for horizon, values in by_horizon.items()},
        "accuracy_flag": accuracy_flag, "calibration_flag": calibration_flag, "verdict": verdict,
    }


def main():
    path = Path(sys.argv[1]) if len(sys.argv) > 1 else Path(__file__).resolve().parent.parent / "datasets" / "msc-p033-weekly-series.csv"
    result = analyze(load(path))
    print(f"origins={len(result['origins'])}")
    print(f"first_origin={result['origins'][0]}")
    print(f"last_origin={result['origins'][-1]}")
    print(f"horizon={HORIZON}")
    print(f"forecasts={result['forecasts']}")
    print(f"model_mase={result['mase']:.6f}")
    print(f"benchmark_mase={result['benchmark']:.6f}")
    print(f"model_mae={result['mae']:.6f}")
    print(f"accuracy_flag={result['accuracy_flag']}")
    for horizon in range(1, HORIZON + 1):
        print(f"mase_h{horizon}={result['by_horizon'][horizon]:.6f}")
    print(f"nominal_coverage={NOMINAL:.2f}")
    print(f"empirical_coverage={result['coverage']:.6f}")
    print(f"mean_interval_width={result['width']:.6f}")
    print(f"mean_interval_score={result['score']:.6f}")
    print(f"calibration_flag={result['calibration_flag']}")
    print(f"verdict={result['verdict']}")


if __name__ == "__main__":
    main()
