import os
import psutil
import threading
import gc
import uuid
import time
import itertools

import numpy as np
import pandas as pd

from sklearn.metrics import adjusted_rand_score, adjusted_mutual_info_score, roc_auc_score, average_precision_score

import json
import sys
import dask
import dask.array as da
import dask.dataframe as dd
from dask.distributed import Client, LocalCluster, get_client, secede, rejoin
from dask_jobqueue import SLURMCluster

from dask_ml.cluster import KMeans as DaskKMeans
from sklearn.ensemble import IsolationForest
from sklearn.cluster import MiniBatchKMeans
from fast_hdbscan import HDBSCAN

import sdobase as sdo
import datagen as dgen

import warnings
warnings.filterwarnings("ignore", category=UserWarning)
warnings.filterwarnings("ignore", category=FutureWarning)

# --- ENVIRONMENT & DASK CONFIGURATIONS ---
os.environ["OMP_NUM_THREADS"] = "1"
os.environ["MKL_NUM_THREADS"] = "1"
os.environ["OPENBLAS_NUM_THREADS"] = "1"
os.environ["DASK_DISTRIBUTED__WORKER__DAEMON"] = "False"

dask.config.set({
    'distributed.comm.timeouts.connect': '90s', # Timeouts for robust networking
    'distributed.comm.timeouts.tcp': '90s',
    'distributed.scheduler.work-steal': False, # Prevents moving tasks if a worker is just "busy"
    'distributed.worker.heartbeat.interval': '5s',
})


def monitor_resources(func):
    def wrapper(*args, **kwargs):
        process = psutil.Process()
        peak_rss = 0
        cpu_usage = []
        running = True

        def monitor():
            nonlocal peak_rss
            while running:
                rss = process.memory_info().rss
                peak_rss = max(peak_rss, rss)
                cpu_percent = psutil.cpu_percent(interval=None)
                cpu_usage.append(cpu_percent)
                time.sleep(0.05)

        t_monitor = threading.Thread(target=monitor)
        t_monitor.start()

        t0 = time.perf_counter()
        res = func(*args, **kwargs)
        t_total = time.perf_counter() - t0

        running = False
        t_monitor.join()

        res["time_sec"] = t_total
        res["peak_mem_mb"] = peak_rss / 1024**2  
        res["cpu_avg_percent"] = np.mean(cpu_usage)
        res["cpu_max_percent"] = np.max(cpu_usage)
        res["num_threads"] = threading.active_count()
        res["num_process_threads"] = process.num_threads()

        client = kwargs.get("client", None)
        if client is not None:
            info = client.scheduler_info()
            total_workers = len(info['workers'])
            total_cores = sum(w.get('ncores', w.get('nthreads', 1)) for w in info['workers'].values())
            total_mem = sum(w.get('memory_limit', 0) for w in info['workers'].values()) / 1024**2
            res.update({
                "dask_workers": total_workers,
                "dask_total_cores": total_cores,
                "dask_total_mem_mb": total_mem
            })

        return res
    return wrapper


class Wp_DaskKMeans:
    def __init__(self, n_clusters, random_state=None):
        self.model = DaskKMeans(n_clusters=n_clusters, random_state=random_state)

    def fit_predict(self, X):
        self.model.fit(X)
        return self.model.predict(X)

class Wp_DaskKMeansAD:
    def __init__(self, n_clusters, random_state=None):
        self.model = DaskKMeans(n_clusters=n_clusters, random_state=random_state)

    def fit(self, X):
        self.model.fit(X)
        return self

    def score_samples(self, X):
        return da.map_blocks(
            lambda block: self.model.transform(block).min(axis=1),
            X, dtype=float, drop_axis=1 )

def dask_cleanup(client=None):
    gc.collect()
    if client is not None:
        client.run(gc.collect)

def build_cluster(cfg):
    if cfg["type"] == "slurm":
        cluster = SLURMCluster(
            name="comparison_tests",
            queue=cfg["queue"],
            cores=cfg["cores"],
            memory=cfg["memory"],
            processes=1,
            walltime=cfg["walltime"],
            log_directory="./results/slurm_logs",      
            local_directory="./results/dask-worker-space",  
            job_extra_directives=cfg["job_extra_directives"]
        )
        cluster.scale(cfg["n_workers"])
    else:
        cluster = LocalCluster(
            n_workers=cfg["n_workers"], threads_per_worker=1, processes=True,
            local_directory="./results/dask-worker-space",  
        )
    return cluster
    
def materialize(arr, array_type):
    result = arr
    if "dask_array" in array_type:
        tmp = sdo.materialize(arr)
        result = np.asarray(tmp).reshape(-1)
    return result

def fit_and_predict(model, X, name, mode="cluster"):
    result = None
    if hasattr(model, "fit_predict"):
        result = model.fit_predict(X)
    else:
        model.fit(X)
        if mode == "cluster":
            result = model.predict(X)
        else:
            result = model.score_samples(X)
    return result


@monitor_resources
def run_clustering_model(model, X, y_true, name, array_type="numpy_array", client=None):

    try:
        y_pred = fit_and_predict(model, X, name, mode="cluster")
        y_pred = materialize(y_pred, array_type)
    except Exception as e:
        return {"task": "clustering", "method": name, "ARI": np.nan, "AMI": np.nan,  "error": f"{type(e).__name__}: {e}"}

    mask = y_true >= 0
    ARI = adjusted_rand_score(y_true[mask], y_pred[mask])
    AMI = adjusted_mutual_info_score(y_true[mask], y_pred[mask])

    return {"task": "clustering", "method": name, "ARI": ARI, "AMI": AMI}

@monitor_resources
def run_ad_model(model, X, y_bin, name, array_type="numpy_array", distance_func=None, client=None):

    try:
        if distance_func:
            model.fit(X)
            scores = distance_func(model, X)
        else:
            scores = fit_and_predict(model, X, name, mode="score")
        
        scores = materialize(scores, array_type)
    except Exception as e:
        return {"task": "anomaly", "method": name, "AUROC": np.nan, "AP": np.nan,  "error": f"{type(e).__name__}: {e}"}

    if name == "IsolationForest":
        scores = -scores

    AUROC = roc_auc_score(y_bin, scores)
    AP = average_precision_score(y_bin, scores)

    return {"task": "anomaly", "method": name, "AUROC": AUROC, "AP": AP}

def target_blocksize(dataset_path, client, tasks_per_worker=4):
    total_bytes = os.path.getsize(dataset_path)
    n_workers = len(client.scheduler_info()['workers'])
    n_partitions = max(n_workers * tasks_per_worker, 1)
    return max(total_bytes // n_partitions, 1)

def load_data(dataset_path, array_type="numpy_array", task="clustering", client=None):
    blocksize = target_blocksize(dataset_path, client) if (client is not None and array_type != "numpy_array") else "200MB"
    dfp = dd.read_parquet(dataset_path, chunksize=blocksize, gather_statistics=False)
    if array_type=="numpy_array":
        X = dfp.drop("label", axis=1).to_dask_array().compute()
    if array_type=="dask_array_no_lenght":
        X = dfp.drop("label", axis=1).to_dask_array().persist()
    if array_type=="dask_array_lenght":
        X = dfp.drop("label", axis=1).to_dask_array(lengths=True).persist()
    y = dfp["label"].compute()
    del dfp
    if task == "anomaly":
        y[y>=0]=0
        y[y<0]=1
    return X,y

def run_clustering(dataset_path, pam):
    results = []
    models = [
        (HDBSCAN(min_cluster_size=10), "fast-HDBSCAN", "numpy_array"),
        (sdo.SDOclust(backend="numpy", n_jobs=-1), "SDOclust-numpy", "numpy_array"),
        (sdo.SDOclust(backend="dask", n_jobs=1), "SDOclust-dask", "dask_array_no_lenght"),
        (MiniBatchKMeans(n_clusters=pam["k"], random_state=pam["rseed"]), "MiniBatchKMeans", "numpy_array"),
        (Wp_DaskKMeans(n_clusters=pam["k"], random_state=pam["rseed"]), "Wp_DaskKMeans", "dask_array_lenght")
        ]
        
    for model, name, array_type in models:
        X,y_true = load_data(dataset_path, array_type, "clustering", client=pam['client'])
        if name=="fast-HDBSCAN" and len(y_true) > 1000000:
            pass
        else:
            res = run_clustering_model(model, X, y_true, name, array_type=array_type,  client=pam['client'])
            del X, y_true
            dask_cleanup(pam["client"])
            print(res)
            results.append(res)
    return results


def run_anomaly_detection(dataset_path, pam):
    results = []

    models = [
        (sdo.SDO(backend="numpy", n_jobs=-1), "SDO-numpy", "numpy_array", None),
        (sdo.SDO(backend="dask", n_jobs=1), "SDO-dask", "dask_array_no_lenght", None),
        (IsolationForest(n_estimators=200, random_state=pam["rseed"], n_jobs=-1), "IsolationForest", "numpy_array", None),
        (Wp_DaskKMeansAD(n_clusters=pam["k"], random_state=pam["rseed"]), "Wp_DaskKMeansAD", "dask_array_lenght", lambda m, X: m.score_samples(X))
        ]

    for model, name, array_type, distance_func in models:
        X,y_true = load_data(dataset_path, array_type, "anomaly", client=pam['client']) 
        res = run_ad_model(model, X, y_true, name, array_type=array_type, distance_func=distance_func, client=pam['client'])
        del X, y_true
        dask_cleanup(pam["client"])
        print(res)
        results.append(res)

    return results


def run_single_experiment(params, client, folder):
    """This function now runs sequentially on the 32GB Manager node."""
    N, d, out, k, option, rseed, fname = params

    # Unique ID for temporary files to prevent workers from overwriting each other
    unique_id = uuid.uuid4().hex[:8]
    dataset_path = f"{folder}/dataset_{unique_id}.parquet"

    # 1. Generate the dataset on the Manager (Plenty of RAM for this)
    X_np, y_np = dgen.generate_dataset(
        N=N, d=d, k=k, outlier_frac=out, seed=rseed,
        option=option, save_path=dataset_path, save_format="parquet"
    )
    del X_np, y_np

    pam = {'k': k, 'rseed': rseed, 'client': client}

    # 2. Run logic (Local models use the Manager, Dask models use the 16 workers!)
    cres = run_clustering(dataset_path, pam)
    ares = run_anomaly_detection(dataset_path, pam)

    # Clean up the worker's local parquet file
    if os.path.exists(dataset_path):
        os.remove(dataset_path)

    combined = cres + ares
    for r in combined:
        r.update({"N": N, "d": d, "k": k, "outlier_frac": out, "dataset_type": option, "rseed": rseed})

    return combined

def run_comparison(cfg, client, fname, folder):
    """Orchestrates the experiments sequentially, leveraging the cluster for math."""
    # Build the grid of all combinations
    param_grid = list(itertools.product(
        cfg["N_list"], cfg["d_list"], cfg["outlier_frac"],
        cfg["k_list"], cfg["option_list"], cfg["rseed"]
    ))

    # Add the output filename to each parameter set
    tasks = [p + (fname,) for p in param_grid]
    print(f"Running {len(tasks)} experiments. (Workers will process the heavy algorithms...)")

    all_results = []

    # Run sequentially to prevent memory explosion
    for task in tasks:
        # Pass the global client into the function
        res = run_single_experiment(task, client, folder)
        all_results.extend(res)

    return pd.DataFrame(all_results)


if __name__ == "__main__":
    folder = "results"
    if not os.path.exists(folder):
        os.mkdir(folder)
        print(f"'{folder}' folder for results created")

    print(f"Results will be saved in '{folder}'...")

    with open(sys.argv[1] if len(sys.argv) > 1 else "config.json") as f:
        run_cfg = json.load(f)

    cluster = build_cluster(run_cfg["cluster"])
    client = Client(cluster)

    print(f"Dask Dashboard link: {client.dashboard_link}")

    print("Waiting for workers to spin up...")
    client.wait_for_workers(n_workers=run_cfg["cluster"]["n_workers"])
    print(f"Connected! Total Workers: {len(client.scheduler_info()['workers'])}")

    experiments = [
        (f"{folder}/results_N.csv", {"N_list": [10000, 100000, 1000000, 10000000], "d_list": [10],
                           "option_list": ["cardinality","density","shapes","groups"], "k_list": [10], "outlier_frac": [0.05], "rseed": [1,2,3]}),
        (f"{folder}/results_d.csv", {"N_list": [100000], "d_list": [5, 10, 50, 100],
                           "option_list": ["cardinality","density","shapes","groups"], "k_list": [10], "outlier_frac": [0.05], "rseed": [1,2,3]}),
        (f"{folder}/results_k.csv", {"N_list": [100000], "d_list": [10],
                           "option_list": ["cardinality","density","shapes","groups"], "k_list": [5, 10, 20, 50], "outlier_frac": [0.05], "rseed": [1,2,3]}),
        (f"{folder}/results_out.csv", {"N_list": [100000], "d_list": [10],
                             "option_list": ["cardinality","density","shapes","groups"], "k_list": [10], "outlier_frac": [0.001, 0.01, 0.05, 0.15], "rseed": [1,2,3]})  ]


    for fname, config in experiments:
        print(f"\n--- Starting {fname} ---")
        print(f"Running experiment: {fname}")

        try:
            df = run_comparison(config, client, fname, folder)
            df.to_csv(fname, index=False)
            print(f"Successfully saved {len(df)} rows to {fname}")
            print(f"Saved results to {fname}")
        except Exception as e:
            print(f"CRITICAL ERROR in {fname}: {e}")
            print(f"Experiment {fname} failed: {e}")

    print("\nAll experiments complete. Shutting down cluster.")
    client.close()
    cluster.close()
