#!/usr/bin/env python3
"""Reproduce the public PharmacyAtlas asking-price alpha metrics."""

from __future__ import annotations

import json
from pathlib import Path

import numpy as np
import pandas as pd
from scipy.stats import spearmanr
from sklearn.linear_model import LinearRegression
from sklearn.metrics import mean_absolute_error, median_absolute_error, r2_score


ROOT = Path(__file__).resolve().parent


def fit_predict(train: pd.DataFrame, test: pd.DataFrame) -> np.ndarray:
    model = LinearRegression().fit(
        np.log1p(train[["annual_prescriptions"]].to_numpy(dtype=float)),
        np.log(train["asking_price_cad"].to_numpy(dtype=float)),
    )
    return np.exp(
        model.predict(
            np.log1p(test[["annual_prescriptions"]].to_numpy(dtype=float))
        )
    )


def metrics(actual: np.ndarray, predicted: np.ndarray) -> dict[str, float]:
    relative_error = np.abs(predicted - actual) / actual
    return {
        "r_squared": float(r2_score(actual, predicted)),
        "mae_cad": float(mean_absolute_error(actual, predicted)),
        "median_absolute_error_cad": float(
            median_absolute_error(actual, predicted)
        ),
        "median_absolute_percentage_error": float(np.median(relative_error)),
        "spearman_rank_correlation": float(
            spearmanr(actual, predicted).statistic
        ),
        "within_25_percent_rate": float(np.mean(relative_error <= 0.25)),
        "within_50_percent_rate": float(np.mean(relative_error <= 0.50)),
    }


def main() -> None:
    frame = pd.read_csv(ROOT / "model-observations.csv")
    actual = frame["asking_price_cad"].to_numpy(dtype=float)

    source_predictions = np.empty(len(frame), dtype=float)
    for source in sorted(frame["marketplace_group"].unique()):
        test_mask = frame["marketplace_group"].eq(source)
        source_predictions[test_mask] = fit_predict(
            frame.loc[~test_mask], frame.loc[test_mask]
        )

    listing_predictions = np.empty(len(frame), dtype=float)
    for index in range(len(frame)):
        listing_predictions[index] = fit_predict(
            frame.drop(frame.index[index]), frame.iloc[[index]]
        )[0]

    result = {
        "rows": int(len(frame)),
        "source_rows": frame["marketplace_group"].value_counts().to_dict(),
        "leave_one_marketplace_out": metrics(actual, source_predictions),
        "leave_one_listing_out": metrics(actual, listing_predictions),
    }
    print(json.dumps(result, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()

