"""
Code developped by Paul Malisani, Adrien Spagnol and Vivien Smis-Michel from IFP Energies nouvelles
"""

import numpy as np
import pickle
import matplotlib.pyplot as plt
import matplotlib


def hyper_parameter_selection():
    matplotlib.rcParams['pdf.fonttype'] = 42
    matplotlib.rcParams['ps.fonttype'] = 42

    with open("hyper_param_selection_datas.pickle", "rb") as fh:
        data_hyper_param = pickle.load(fh)

    # load results for rpha with varying alpha results
    dict_rpha = [d for d in data_hyper_param if d["method"] == "rpha"]
    for i in range(len(dict_rpha) - 1):
        assert dict_rpha[i]["rolling_horizon_in_hours"] == dict_rpha[i+1]["rolling_horizon_in_hours"]
        assert dict_rpha[i]["Ns"] == dict_rpha[i + 1]["Ns"]
        assert dict_rpha[i]["Nred"] == dict_rpha[i + 1]["Nred"]

    Ns, Nred, rolling_horizon = dict_rpha[0]["Ns"], dict_rpha[0]["Nred"], dict_rpha[0]["rolling_horizon_in_hours"]

    # retrieve couples rpha weight - performance ratio
    rpha_weights = [d["rpha_weight"] for d in dict_rpha]
    performance_ratios = [d["performance_ratio"] for d in dict_rpha]


    fontsize_xy_label = 24
    fontsize_xy_ticks = 24
    fontsize_title = 24

    fig, ax = plt.subplots()  # Use subplots() instead of figure()

    ax.plot(
        rpha_weights, performance_ratios,
        linestyle="dashed", linewidth=5, marker="o", markersize=15,
        label="$N_{s}$ = " + str(Ns) + " ; $N_{\\rm red}$ = " + str(Nred)
    )

    # Axis labels
    ax.set_ylabel(
        "Performance ratio $\eta(N_{\\rm red}, \\alpha)$",
        fontdict=dict(size=fontsize_xy_label)
    )
    ax.set_xlabel(
        "VPPHA weighting parameter $\\alpha$",
        fontdict=dict(size=fontsize_xy_label)
    )

    # Tick font sizes
    ax.tick_params(axis='x', labelsize=fontsize_xy_ticks)
    ax.tick_params(axis='y', labelsize=fontsize_xy_ticks)


    # Title
    ax.set_title(
        "Performance ratio $\eta(N_{\\rm red}, \\alpha)$ with a rolling-horizon period $H = $ "
        + str(rolling_horizon) + " hours",
        fontdict=dict(size=fontsize_title)
    )

    # Optional: legend and layout adjustment
    #ax.legend()
    plt.tight_layout()



def two_years_simulation():
    matplotlib.rcParams['pdf.fonttype'] = 42
    matplotlib.rcParams['ps.fonttype'] = 42
    with open("two_years_simulation.pickle", "rb") as fh:
        data_two_years = pickle.load(fh)

    fontsize_xy_label = 24
    fontsize_xy_ticks = 24
    fontsize_legend = 24
    fontsize_title = 24

    dict_mpc, dict_std_pha, dict_rpha = data_two_years["mpc"], data_two_years["std_pha"], data_two_years["rpha"]

    # --- Figure 1: Performance ratio evolution ---
    fig1, ax1 = plt.subplots()

    ax1.plot(
        dict_std_pha["time_date"], dict_std_pha["evolution_performance_ratio"],
        label="Standard PHA", linewidth=3, color='tab:blue'
    )
    ax1.plot(
        dict_rpha["time_date"], dict_rpha["evolution_performance_ratio"],
        label="VPPHA", linewidth=3, color='tab:red'
    )
    ax1.plot(
        dict_rpha["time_date"], np.zeros_like(dict_rpha["time"]),
        linewidth=2, linestyle="--", color='black'
    )

    # Axis labels and title
    ax1.set_xlabel("Date", fontdict=dict(size=fontsize_xy_label))
    ax1.set_ylabel("Performance ratio $\eta$ ; $H=24$", fontdict=dict(size=fontsize_xy_label))
    ax1.set_title("Time evolution of the performance ratio $\eta(\\alpha)$",
                  fontdict=dict(size=fontsize_title))

    # Ticks and limits
    ax1.tick_params(axis='x', labelsize=fontsize_xy_ticks)
    ax1.tick_params(axis='y', labelsize=fontsize_xy_ticks)
    ax1.set_ylim(-20., 25.)

    # Optional: auto-format x-axis dates
    fig1.autofmt_xdate()
    # If you want fewer date ticks:
    # ax1.xaxis.set_major_locator(mdates.AutoDateLocator(maxticks=8))

    ax1.legend(prop={'size': fontsize_legend})
    plt.tight_layout()

    # --- Figure 2: Bill reduction ---
    fig2, ax2 = plt.subplots()

    ax2.plot(
        dict_mpc["time_date"], dict_mpc["evolution_bill_reduction"],
        label="MPC", linewidth=3, color='tab:green'
    )
    ax2.plot(
        dict_std_pha["time_date"], dict_std_pha["evolution_bill_reduction"],
        label="Standard PHA", linewidth=3, color='tab:blue'
    )
    ax2.plot(
        dict_rpha["time_date"], dict_rpha["evolution_bill_reduction"],
        label="VPPHA", linewidth=3, color='tab:red'
    )

    # Axis labels and title
    ax2.set_xlabel("Date", fontdict=dict(size=fontsize_xy_label))
    ax2.set_ylabel("Percentage of bill reduction",
                   fontdict=dict(size=fontsize_xy_label))
    ax2.set_title("Electricity bill reduction with respect to storage-less battery bill",
                  fontdict=dict(size=fontsize_title))

    # Ticks and limits
    ax2.tick_params(axis='x', labelsize=fontsize_xy_ticks)
    ax2.tick_params(axis='y', labelsize=fontsize_xy_ticks)
    ax2.set_ylim(0., 10.)

    # Auto-format x-axis dates
    fig2.autofmt_xdate()
    # ax2.xaxis.set_major_locator(mdates.AutoDateLocator(maxticks=8))

    ax2.legend(prop={'size': fontsize_legend})
    plt.tight_layout()

    print(" ")
    print("Performance ratio at the end of simulation for standard PHA = ", dict_std_pha["evolution_performance_ratio"][-1])
    print("Performance ratio at the end of simulation for VPPHA = ", dict_rpha["evolution_performance_ratio"][-1])
    print(" ")
    print("Bill reduction at the end of simulation for MPC = ", dict_mpc["evolution_bill_reduction"][-1])
    print("Bill reduction at the end of simulation for standard PHA = ", dict_std_pha["evolution_bill_reduction"][-1])
    print("Bill reduction at the end of simulation for VPPHA = ", dict_rpha["evolution_bill_reduction"][-1])


if __name__ == "__main__":
    hyper_parameter_selection()
    two_years_simulation()
    plt.show()