"""MSC-P-022 associative price-elasticity reference. Original MIT example."""
from __future__ import annotations

import csv
import math
from pathlib import Path

DATA = Path(__file__).parents[1] / "datasets" / "msc-001-price-elasticity.csv"


def solve(matrix: list[list[float]], vector: list[float]) -> list[float]:
    augmented = [row[:] + [value] for row, value in zip(matrix, vector)]
    size = len(vector)
    for column in range(size):
        pivot = max(range(column, size), key=lambda row: abs(augmented[row][column]))
        if abs(augmented[pivot][column]) < 1e-12:
            raise ValueError("Singular design matrix")
        augmented[column], augmented[pivot] = augmented[pivot], augmented[column]
        scale = augmented[column][column]
        augmented[column] = [value / scale for value in augmented[column]]
        for row in range(size):
            if row == column:
                continue
            factor = augmented[row][column]
            augmented[row] = [left - factor * right for left, right in zip(augmented[row], augmented[column])]
    return [augmented[row][-1] for row in range(size)]


def inverse(matrix: list[list[float]]) -> list[list[float]]:
    size = len(matrix)
    columns = []
    for index in range(size):
        unit = [0.0] * size
        unit[index] = 1.0
        columns.append(solve(matrix, unit))
    return [[columns[column][row] for column in range(size)] for row in range(size)]


def ols(design: list[list[float]], outcome: list[float]) -> tuple[list[float], list[float]]:
    width = len(design[0])
    xtx = [[sum(row[i] * row[j] for row in design) for j in range(width)] for i in range(width)]
    xty = [sum(row[i] * value for row, value in zip(design, outcome)) for i in range(width)]
    beta = solve(xtx, xty)
    residuals = [value - sum(coef * item for coef, item in zip(beta, row)) for row, value in zip(design, outcome)]
    variance = sum(value * value for value in residuals) / (len(outcome) - width)
    covariance = inverse(xtx)
    standard_errors = [math.sqrt(variance * covariance[i][i]) for i in range(width)]
    return beta, standard_errors


def main() -> None:
    with DATA.open(encoding="utf-8", newline="") as stream:
        rows = list(csv.DictReader(stream))
    if len(rows) != 24 or any(float(row["price_eur"]) <= 0 or float(row["units"]) <= 0 for row in rows):
        raise ValueError("Expected 24 strictly positive price and quantity observations")
    y = [math.log(float(row["units"])) for row in rows]
    simple_x = [[1.0, math.log(float(row["price_eur"]))] for row in rows]
    controlled_x = [[1.0, math.log(float(row["price_eur"])), float(row["season_index"])] for row in rows]
    simple_beta, simple_se = ols(simple_x, y)
    controlled_beta, controlled_se = ols(controlled_x, y)
    elasticity, se = controlled_beta[1], controlled_se[1]
    low, high = elasticity - 1.959963984540054 * se, elasticity + 1.959963984540054 * se
    exact_volume_change = math.pow(1.05, elasticity) - 1
    print(f"simple_elasticity={simple_beta[1]:.6f} se={simple_se[1]:.6f}")
    print(f"season_controlled_elasticity={elasticity:.6f} se={se:.6f}")
    print(f"ci95=[{low:.6f}, {high:.6f}]")
    print(f"conditional_volume_change_at_plus_5pct={100 * exact_volume_change:.6f}%")


if __name__ == "__main__":
    main()
