
import os
import numpy as np
import warnings
np.warnings = warnings
import itertools

from gen_data import gen_data

from sklearn.preprocessing import StandardScaler
from sklearn.model_selection import train_test_split
from sklearn.metrics import (roc_auc_score, average_precision_score, adjusted_rand_score, adjusted_mutual_info_score)

from pyod.models.loda    import LODA
from pyod.models.abod    import ABOD
from pyod.models.hbos    import HBOS
from pyod.models.knn     import KNN
from pyod.models.iforest import IForest
from pyod.models.lof     import LOF
from sklearn.cluster     import KMeans, MiniBatchKMeans
from pyod.models.ecod    import ECOD
from pyod.models.copod   import COPOD
import hdbscan
from classix import CLASSIX
import sdoclust as sdo
from scipy.cluster.hierarchy import fcluster
import fastcluster
from clustpy.partition import SkinnyDip  
from pyclustering.cluster.xmeans import xmeans
from pyclustering.cluster.fcm import fcm


def list_local_datasets(datasets_dir):
    if not os.path.isdir(datasets_dir):
        raise FileNotFoundError(f"Non-existing dataset folder: {datasets_dir}")
    names = sorted(f for f in os.listdir(datasets_dir) if f.endswith(".npz"))
    if not names:
        raise FileNotFoundError(f"No .npz files found in {datasets_dir}")
    return names


def load_dataset(name, datasets_dir="datasets/Classical", max_samples=10000, seed=0, scale=True):
    path = os.path.join(datasets_dir, name)
    if not os.path.exists(path):
        raise FileNotFoundError(f"No dataset found: {path}")
    data = np.load(path, allow_pickle=True)
    X, y = data["X"].astype(np.float64), data["y"].astype(int)

    if max_samples is not None and len(y) > max_samples:
        # Subsample while preserving the original class ratio.
        X, _, y, _ = train_test_split( X, y, train_size=max_samples, stratify=y, random_state=seed )

    if scale:
        X = StandardScaler().fit_transform(X)
        
    # Keep contamination within the range expected by the detectors.
    contamination = float(np.clip(y.mean(), 0.01, 0.5))
    return X, y, contamination


def generate_data(n_samples, n_features, n_clusters, outlier_fraction, scenario="easy", seed=0):

    if scenario=="easy":
        difficulty = 0.2
    else:
        difficulty = 0.8
    X_in, y_in, X, yb, yc = gen_data(n_samples, n_features, n_clusters, outlier_fraction, difficulty=difficulty, random_state=seed)

    return X_in, y_in, X, yb, yc

def generate_interaction_configurations(scenario="easy", fractional=False):
    """2-level factorial design at axis extremes, to test interactions
    beyond the OFAT screening already covered by generate_dataset_configurations.
    fractional=True uses a resolution-IV half-fraction (2^(4-1)) to avoid
    the most expensive corner (all factors at max simultaneously)."""
    factor = 1000
    levels = {
        "n_samples":        [5*factor, 5000*factor],
        "n_features":       [5, 1000],
        "n_clusters":       [2, 200],
        "outlier_fraction": [0.01, 0.40],
    }
    factors = list(levels.keys())
    if not fractional:
        combos = list(itertools.product(*levels.values()))
    else:
        # Resolution IV half-fraction: (max,max,max,max) is never generated.
        base_signs = list(itertools.product([-1, 1], repeat=3))
        combos = []
        for s1, s2, s3 in base_signs:
            s4 = -s1 * s2 * s3
            combos.append(tuple(
                levels[f][0] if s == -1 else levels[f][1]
                for f, s in zip(factors, (s1, s2, s3, s4))   ))
    return combos  
    
def generate_dataset_configurations(analysis_type, scenario="easy"):
    """Build the list of (n_samples, n_features, n_clusters, outlier_fraction) configs.
    Only one axis varies per analysis_type; the rest are fixed at baseline values."""

    factor = 1000
    n_samples_list        = [5, 10, 50, 100, 500, 5000]
    n_samples_list        = [n * factor for n in n_samples_list]
    n_features_list       = [5, 10, 50, 100, 500, 1000]
    n_clusters_list       = [2, 5, 10, 50, 100, 200]
    outlier_fraction_list = [0.01, 0.05, 0.1, 0.2, 0.3, 0.4]

    # Baseline values (fixed axes)
    n = len(n_samples_list)
    n_samples        = [100000] * n
    n_features       = [10]     * n
    n_clusters       = [10]     * n
    outlier_fraction = [0.05]   * n

    # Vary only the relevant axis
    if   analysis_type == "size": n_samples        = n_samples_list
    elif analysis_type == "dims": n_features       = n_features_list
    elif analysis_type == "clus": n_clusters       = n_clusters_list
    elif analysis_type == "outs": outlier_fraction = outlier_fraction_list

    return [(n_samples[i], n_features[i], n_clusters[i], outlier_fraction[i]) for i in range(n)]


class FastHierarchicalWrapper:

    def __init__(self, n_clusters):
        self.n_clusters = n_clusters
        self.labels_    = None

    def fit(self, X):
        # Build the linkage tree, then cut it into n_clusters groups.
        Z            = fastcluster.linkage_vector(X, method='ward')
        self.labels_ = fcluster(Z, t=self.n_clusters, criterion='maxclust')
        return self

    def fit_predict(self, X):
        return self.fit(X).labels_


class XMeansWrapper:
    def __init__(self, max_n_clusters=20, random_state=None):
        self.max_n_clusters = max_n_clusters
        self.random_state = random_state

    def fit(self, X):
        self.model = xmeans( X.tolist(),  kmax=self.max_n_clusters,  ccore=True, random_state=self.random_state  )
        self.model.process()

        self.labels_ = np.full(len(X), -1, dtype=int)
        for i, cluster in enumerate(self.model.get_clusters()):
            # Convert pyclustering's index lists into sklearn-style labels.
            self.labels_[cluster] = i

        return self

    def fit_predict(self, X):
        return self.fit(X).labels_


class FCMWrapper:
    def __init__(self, n_clusters=2, random_state=None):
        self.n_clusters = n_clusters
        self.random_state = random_state

    def fit(self, X):
        rng = np.random.default_rng(self.random_state)
        # pyclustering expects initial cluster centers rather than n_clusters.
        centers = X[rng.choice(len(X), self.n_clusters, replace=False)]

        self.model = fcm(X.tolist(), centers.tolist(), ccore=True)
        self.model.process()
        self.labels_ = np.argmax(self.model.get_membership(), axis=1)

        return self

    def fit_predict(self, X):
        return self.fit(X).labels_        
        
        
def run_ad(config, Xo, yo):
    model    = config["model"](**config["params"])
    preds    = model.fit_predict(Xo)
    return roc_auc_score(yo, preds), average_precision_score(yo, preds)


def run_cl(name, config, Xo, yc):
    model = config["model"](**config["params"])
    try:
        preds = model.fit_predict(Xo)
    except Exception:
        # Some sklearn estimators don't implement fit_predict; fall back to fit + labels_
        model.fit(Xo)
        preds = model.labels_

    # Exclude points labeled as anomalies in the ground-truth
    valid = (yc >= 0)
    if valid.sum() > 0:
        ari = adjusted_rand_score(yc[valid], preds[valid])
        ami = adjusted_mutual_info_score(yc[valid], preds[valid])
    else:
        ari = ami = np.nan
    return ari, ami
    
def n_bins_scaled(size, b_min=10, b_max=100, b_ref=20, size_ref=1000):
    b = b_ref + 10 * np.log10(max(size, 1) / size_ref)
    return int(np.clip(round(b), b_min, b_max))
    
def neighborhood_k(size, k_min=10, k_max=100, k_ref=20, size_ref=1000):
    k = k_ref + 5 * np.log10(max(size, 1) / size_ref)
    return int(np.clip(round(k), k_min, k_max))
    
def set_analysis_algorithms(n_clusters, size, dims, n_jobs=1, seed=0):
    k_nb       = neighborhood_k(size)
    n_cuts_loda = min(150, int(np.ceil(np.sqrt(size))))  
    n_bins = n_bins_scaled(size)

    ad_algorithms = {
        "ad_SDOchunk": {"model": sdo.SDO,  "params": {"chunksize": 10000}},
        "ad_LODA":     {"model": LODA,     "params": {"n_bins": n_bins, "n_random_cuts": n_cuts_loda}},
        "ad_FastABOD": {"model": ABOD,     "params": {"method": "fast", "n_neighbors": 5}},
        "ad_HBOS":     {"model": HBOS,     "params": {"n_bins": n_bins}},
        "ad_kNN":      {"model": KNN,      "params": {"method": "largest", "n_jobs": n_jobs, "n_neighbors": k_nb, "algorithm": "auto"}},
        "ad_IForest":  {"model": IForest,  "params": {"random_state": seed, "n_jobs": n_jobs, "n_estimators": 100, "max_samples": 256}},
        "ad_LOF":      {"model": LOF,      "params": {"n_neighbors": k_nb, "n_jobs": n_jobs, "algorithm": "auto"}},
        "ad_ECOD":     {"model": ECOD,     "params": {}},
        "ad_COPOD":    {"model": COPOD,    "params": {}} }

    common     = {"chi": 5, "k": 1000} if n_clusters > 10 else {}
    clx_radius = 0.15 * (100000 / size) ** 0.2
    clx_radius = np.clip(clx_radius, 0.1, 0.5)
    clx_minPts = max(5, round(5 * (size / 10000) ** np.log10(2)))

    clu_algorithms = {
        "cl_SDOCLchunk":       {"model": sdo.SDOclust,            "params": {**common, "chunksize": 10000}},
        "cl_CLASSIX":          {"model": CLASSIX,                 "params": {"verbose": 0, "minPts": clx_minPts, "radius": clx_radius}},
        "cl_HDBSCAN":          {"model": hdbscan.HDBSCAN,         "params": {"min_cluster_size": 5, "core_dist_n_jobs": n_jobs, "approx_min_span_tree": True}},
        "cl_KMeans":           {"model": KMeans,                  "params": {"n_clusters": n_clusters, "random_state": seed, "n_init": "auto"}},
        "cl_MBkMeans":         {"model": MiniBatchKMeans,         "params": {"n_clusters": n_clusters, "random_state": seed}},
        "cl_FastHierarchical": {"model": FastHierarchicalWrapper, "params": {"n_clusters": n_clusters}},
        "cl_XMeans":           {"model": XMeansWrapper,           "params": {"max_n_clusters": n_clusters * 3, "random_state": seed}},
        "cl_SkinnyDip":        {"model": SkinnyDip,               
            "params": {"significance": 0.05, "pval_strategy": "table", "outliers": True, "add_tails": True, "random_state": seed}},
        "cl_FCM":              {"model": FCMWrapper,              "params": {"n_clusters": n_clusters, "random_state": seed}} }
        
    return ad_algorithms, clu_algorithms
