import os

os.environ["OMP_NUM_THREADS"] = "1"
os.environ["MKL_NUM_THREADS"] = "1"
os.environ["OPENBLAS_NUM_THREADS"] = "1"

import time
import pandas as pd
import sys
import psutil
import numpy as np

import datagen as dgen
import sdobase as sdo

from sklearn.metrics import adjusted_rand_score, adjusted_mutual_info_score
from stream_test import load_config, build_client


def random_sample_dask(X, n, random_state=0):
    if not isinstance(X.shape[0], (int, np.integer)):
        raise ValueError("X must have known length (lengths=True)")

    rng = np.random.RandomState(random_state)
    m = X.shape[0]
    idx = rng.choice(m, size=min(n, m), replace=False)
    return X[idx]

def get_worker_rss():
    return psutil.Process().memory_info().rss

if __name__ == "__main__":

    if len(sys.argv) < 3:
        print("Usage: python scaling.py <type_of_scaling> <execution.json>")
        sys.exit(1)

    scaling = sys.argv[1]   # "weak" or "strong"
    execution_config = load_config(sys.argv[2])
    base_cluster_config = execution_config.get("cluster", {})

    folder = "results"
    try:
        os.mkdir(folder)
        print(f"'{folder}' folder for results created")
    except FileExistsError:
        print(f"Results will be saved in '{folder}'...")


    N_ref = 100000
    d = 25
    k = 15
    outlier_frac = 0.01
    dataset_type = "cardinality"
    rseed = 42

    workers_list = [1, 2, 4, 8, 16]
    results = []

    for n_workers in workers_list:

        if scaling == "weak":
            print(f"\n### WEAK SCALING: {n_workers} workers ###")
            N_total = N_ref * n_workers
        else:
            print(f"\n### STRONG SCALING: {n_workers} workers ###")
            N_total = N_ref * 10

        N_per_worker = N_total // n_workers

        dataset_path = f"{folder}/dataset_{scaling}_{N_total}.parquet"

        print("Generating dataset...")
        X_np, y_np = dgen.generate_dataset( N=N_total, d=d, k=k, outlier_frac=outlier_frac, seed=rseed,
            option=dataset_type, save_path=dataset_path, save_format="parquet" )
        del X_np, y_np

        # Same infrastructure (queue/cores/memory/walltime/...) as the
        # execution.json, only n_workers changes across the scaling sweep.
        cluster_config = dict(base_cluster_config)
        cluster_config["n_workers"] = n_workers
        client, cluster = build_client(cluster_config)

        client.wait_for_workers(n_workers=n_workers)
        n_workers_eff = len(client.scheduler_info()['workers'])
        assert n_workers_eff == n_workers, f"expected {n_workers} workers, got {n_workers_eff}"

        print(client)
        print(client.nthreads())

        import dask.dataframe as dd

        df = dd.read_parquet(dataset_path)
        X_dask = df.drop("label", axis=1).to_dask_array(lengths=True)
        y = df["label"].compute()
        del df

        tasks_per_worker = 1 if scaling == "weak" else 4
        n_partitions = n_workers_eff * tasks_per_worker
        rows_per_partition = max(X_dask.shape[0] // n_partitions, 1)
        X_dask = X_dask.rechunk((rows_per_partition, X_dask.shape[1]))

        Xs = random_sample_dask(X_dask, N_ref, random_state=rseed)
        Xs = sdo.materialize(Xs)

        t0 = time.perf_counter()
        model = sdo.SDOclust(backend="dask", n_jobs=1).fit(Xs)
        t_fit = time.perf_counter() - t0

        t1 = time.perf_counter()
        labels = model.predict(X_dask)
        labels = sdo.materialize(labels)
        t_predict = time.perf_counter() - t1

        worker_mem = client.run(get_worker_rss)
        peak_mem_mb = max(worker_mem.values()) / 1024**2

        ari = adjusted_rand_score(y[y >= 0], labels[y >= 0])
        ami = adjusted_mutual_info_score(y[y >= 0], labels[y >= 0])
        
        del labels, y, Xs
        
        res = {"scaling": scaling, "workers": n_workers, "workers_eff": n_workers_eff, "n_partitions": X_dask.npartitions,
            "N": N_total, "N_per_worker": N_per_worker,
            "time_fit": t_fit, "time_predict": t_predict, "peak_mem_mb": peak_mem_mb, "ARI": ari, "AMI": ami }

        print(res)
        results.append(res)

        client.close()
        if cluster is not None:
            cluster.close()

        if os.path.exists(dataset_path):
            os.remove(dataset_path)

    df_res = pd.DataFrame(results)

    T1 = df_res.loc[df_res["workers"] == 1, "time_predict"].values[0]

    if scaling == "strong":
        df_res["speedup"] = T1 / df_res["time_predict"]
        df_res["efficiency"] = df_res["speedup"] / df_res["workers"]
    else:
        df_res["weak_efficiency"] = T1 / df_res["time_predict"]

    df_res.to_csv(f"{folder}/{scaling}_scaling.csv", index=False)
