#!/usr/bin/env python3
"""Reproduce the synthetic MSC-P-018 predictive-versus-causal example.

Source code: MIT. Generated dataset: CC0 1.0.
The data are synthetic and describe no real campaign.
"""

from __future__ import annotations

import argparse
import csv
import math
import random
from pathlib import Path

SEED = 20260819
N = 2400
TRAIN_N = 1800
TRUE_DIRECT_EFFECT = 2.5
TRUE_MEDIATOR_EFFECT = 4.5
TRUE_TREATMENT_TO_MEDIATOR = 1.4
TRUE_TOTAL_EFFECT = TRUE_DIRECT_EFFECT + TRUE_MEDIATOR_EFFECT * TRUE_TREATMENT_TO_MEDIATOR


def transpose(matrix: list[list[float]]) -> list[list[float]]:
    return [list(column) for column in zip(*matrix)]


def matmul(left: list[list[float]], right: list[list[float]]) -> list[list[float]]:
    return [[sum(a * b for a, b in zip(row, column)) for column in transpose(right)] for row in left]


def invert(matrix: list[list[float]]) -> list[list[float]]:
    n = len(matrix)
    augmented = [row[:] + [1.0 if i == j else 0.0 for j in range(n)] for i, row in enumerate(matrix)]
    for column in range(n):
        pivot = max(range(column, n), 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(n):
            if row == column:
                continue
            factor = augmented[row][column]
            augmented[row] = [a - factor * b for a, b in zip(augmented[row], augmented[column])]
    return [row[n:] for row in augmented]


def ols(x: list[list[float]], y: list[float]) -> tuple[list[float], list[float], list[float]]:
    xt = transpose(x)
    bread = invert(matmul(xt, x))
    beta = [row[0] for row in matmul(matmul(bread, xt), [[value] for value in y])]
    residuals = [actual - sum(coef * value for coef, value in zip(beta, row)) for row, actual in zip(x, y)]
    n, p = len(x), len(x[0])
    meat = [[0.0 for _ in range(p)] for _ in range(p)]
    for row, residual in zip(x, residuals):
        for j in range(p):
            for k in range(p):
                meat[j][k] += residual * residual * row[j] * row[k]
    covariance = matmul(matmul(bread, meat), bread)
    hc1 = n / (n - p)
    robust_se = [math.sqrt(max(0.0, hc1 * covariance[j][j])) for j in range(p)]
    return beta, robust_se, residuals


def predict(beta: list[float], x: list[list[float]]) -> list[float]:
    return [sum(coef * value for coef, value in zip(beta, row)) for row in x]


def metrics(actual: list[float], predicted: list[float]) -> tuple[float, float, float]:
    errors = [a - p for a, p in zip(actual, predicted)]
    rmse = math.sqrt(sum(error * error for error in errors) / len(errors))
    mae = sum(abs(error) for error in errors) / len(errors)
    mean_y = sum(actual) / len(actual)
    r2 = 1.0 - sum(error * error for error in errors) / sum((value - mean_y) ** 2 for value in actual)
    return rmse, mae, r2


def paired_rmse_improvement_interval(actual: list[float], pre: list[float], post: list[float], draws: int = 2000) -> tuple[float, float, float]:
    observed = 100.0 * (1.0 - metrics(actual, post)[0] / metrics(actual, pre)[0])
    rng = random.Random(SEED + 1)
    bootstrap: list[float] = []
    for _ in range(draws):
        indices = [rng.randrange(len(actual)) for _ in actual]
        sampled_actual = [actual[index] for index in indices]
        sampled_pre = [pre[index] for index in indices]
        sampled_post = [post[index] for index in indices]
        bootstrap.append(100.0 * (1.0 - metrics(sampled_actual, sampled_post)[0] / metrics(sampled_actual, sampled_pre)[0]))
    bootstrap.sort()
    return observed, bootstrap[int(0.025 * draws)], bootstrap[int(0.975 * draws) - 1]


def generate_rows() -> list[dict[str, float | int | str]]:
    rng = random.Random(SEED)
    rows: list[dict[str, float | int | str]] = []
    for index in range(N):
        intent = rng.gauss(0.0, 1.0)
        prior_spend = 50.0 + 12.0 * rng.gauss(0.0, 1.0)
        treatment = int(rng.random() < 0.5)
        engagement = 0.5 + TRUE_TREATMENT_TO_MEDIATOR * treatment + 0.9 * intent + 0.015 * prior_spend + rng.gauss(0.0, 0.8)
        spend = 25.0 + TRUE_DIRECT_EFFECT * treatment + TRUE_MEDIATOR_EFFECT * engagement + 3.0 * intent + 0.12 * prior_spend + rng.gauss(0.0, 5.0)
        rows.append({
            "customer_id": f"C{index + 1:04d}",
            "split": "train" if index < TRAIN_N else "holdout",
            "baseline_intent_z": intent,
            "prior_spend_eur": prior_spend,
            "reminder_assigned": treatment,
            "engagement_index_7d": engagement,
            "spend_30d_eur": spend,
        })
    return rows


def write_csv(path: Path, rows: list[dict[str, float | int | str]]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    fields = ["customer_id", "split", "baseline_intent_z", "prior_spend_eur", "reminder_assigned", "engagement_index_7d", "spend_30d_eur"]
    with path.open("w", encoding="utf-8", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=fields, lineterminator="\n")
        writer.writeheader()
        for row in rows:
            writer.writerow({key: f"{value:.6f}" if isinstance(value, float) else value for key, value in row.items()})


def read_csv(path: Path) -> list[dict[str, float | int | str]]:
    with path.open(encoding="utf-8", newline="") as handle:
        rows = list(csv.DictReader(handle))
    return [{
        "customer_id": row["customer_id"],
        "split": row["split"],
        "baseline_intent_z": float(row["baseline_intent_z"]),
        "prior_spend_eur": float(row["prior_spend_eur"]),
        "reminder_assigned": int(row["reminder_assigned"]),
        "engagement_index_7d": float(row["engagement_index_7d"]),
        "spend_30d_eur": float(row["spend_30d_eur"]),
    } for row in rows]


def design(rows: list[dict[str, float | int | str]], include_mediator: bool) -> list[list[float]]:
    matrix = [[1.0, float(row["baseline_intent_z"]), float(row["prior_spend_eur"]), float(row["reminder_assigned"])] for row in rows]
    if include_mediator:
        for values, row in zip(matrix, rows):
            values.append(float(row["engagement_index_7d"]))
    return matrix


def main() -> None:
    parser = argparse.ArgumentParser()
    default_csv = Path(__file__).resolve().parents[1] / "datasets" / "msc-p018-predictive-causal.csv"
    parser.add_argument("--csv", type=Path, default=default_csv)
    parser.add_argument("--generate", action="store_true", help="regenerate the deterministic synthetic CSV")
    args = parser.parse_args()

    if args.generate or not args.csv.exists():
        write_csv(args.csv, generate_rows())
    rows = read_csv(args.csv)
    train = [row for row in rows if row["split"] == "train"]
    holdout = [row for row in rows if row["split"] == "holdout"]
    y_train = [float(row["spend_30d_eur"]) for row in train]
    y_holdout = [float(row["spend_30d_eur"]) for row in holdout]
    y_all = [float(row["spend_30d_eur"]) for row in rows]

    predictive_pre, _, _ = ols(design(train, False), y_train)
    predictive_post, _, _ = ols(design(train, True), y_train)
    pre_predictions = predict(predictive_pre, design(holdout, False))
    post_predictions = predict(predictive_post, design(holdout, True))
    pre_metrics = metrics(y_holdout, pre_predictions)
    post_metrics = metrics(y_holdout, post_predictions)
    improvement = paired_rmse_improvement_interval(y_holdout, pre_predictions, post_predictions)

    total_beta, total_se, _ = ols(design(rows, False), y_all)
    direct_beta, direct_se, _ = ols(design(rows, True), y_all)
    total_low = total_beta[3] - 1.96 * total_se[3]
    total_high = total_beta[3] + 1.96 * total_se[3]
    direct_low = direct_beta[3] - 1.96 * direct_se[3]
    direct_high = direct_beta[3] + 1.96 * direct_se[3]

    print(f"seed: {SEED}")
    print(f"rows: {len(rows)} (train={len(train)}, holdout={len(holdout)})")
    print(f"true_total_effect_eur: {TRUE_TOTAL_EFFECT:.6f}")
    print(f"predictive_pre_rmse_eur: {pre_metrics[0]:.6f}")
    print(f"predictive_pre_mae_eur: {pre_metrics[1]:.6f}")
    print(f"predictive_pre_r2: {pre_metrics[2]:.6f}")
    print(f"predictive_post_rmse_eur: {post_metrics[0]:.6f}")
    print(f"predictive_post_mae_eur: {post_metrics[1]:.6f}")
    print(f"predictive_post_r2: {post_metrics[2]:.6f}")
    print(f"predictive_rmse_improvement_pct: {improvement[0]:.6f}")
    print(f"predictive_rmse_improvement_bootstrap_95ci: [{improvement[1]:.6f}, {improvement[2]:.6f}]")
    print(f"causal_total_ate_eur: {total_beta[3]:.6f}")
    print(f"causal_total_hc1_se: {total_se[3]:.6f}")
    print(f"causal_total_95ci: [{total_low:.6f}, {total_high:.6f}]")
    print(f"post_treatment_adjusted_treatment_coef_eur: {direct_beta[3]:.6f}")
    print(f"post_treatment_adjusted_hc1_se: {direct_se[3]:.6f}")
    print(f"post_treatment_adjusted_95ci: [{direct_low:.6f}, {direct_high:.6f}]")
    print("scope: prediction metrics are holdout estimates; causal coefficient uses the full randomized synthetic sample")


if __name__ == "__main__":
    main()
