import gc
import json
import os
import sys
import threading
import time
from collections import deque

import numpy as np
import pandas as pd
import psutil

import dask
import dask.array as da
import pyarrow.dataset as ds
import pyarrow.compute as pc
from dask.distributed import Client
from dask_jobqueue import SLURMCluster
from dask_ml.cluster import KMeans as DaskKMeans

from sklearn.metrics import adjusted_rand_score, adjusted_mutual_info_score
from sklearn.cluster import MiniBatchKMeans
from sklearn.ensemble import IsolationForest
from fast_hdbscan import HDBSCAN

import sdobase as sdo
import fast_metrics as fm

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

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

dask.config.set({
    'distributed.comm.timeouts.connect': '90s',
    'distributed.comm.timeouts.tcp': '90s',
    'distributed.scheduler.work-steal': False,
    'distributed.worker.heartbeat.interval': '5s',
})


def load_config(config_path):
    with open(config_path) as f:
        config = json.load(f)
    return config


def get_file_boundaries(parquet_files):
    file_boundaries = []
    row_offset = 0
    for path in parquet_files:
        dataset = ds.dataset(path, format="parquet")
        n_rows = dataset.count_rows()  # metadata only, no data materialized
        file_boundaries.append((os.path.basename(path), row_offset, row_offset + n_rows))
        row_offset += n_rows
    return file_boundaries, row_offset


def load_targets(parquet_files, target, binarize_target_0class=None, read_batch_rows=1_000_000):
    chunks = []
    for path in parquet_files:
        dataset = ds.dataset(path, format="parquet")
        scanner = dataset.scanner(columns=[target], batch_size=read_batch_rows)
        for record_batch in scanner.to_batches():
            col = record_batch.column(0)
            if binarize_target_0class is not None:
                is_zero_class = pc.equal(col, binarize_target_0class)
                y_chunk = np.where(is_zero_class.to_numpy(zero_copy_only=False), 0, 1).astype(np.int8)
            else:
                y_chunk = col.to_numpy(zero_copy_only=False)
            chunks.append(y_chunk)
    return np.concatenate(chunks) if chunks else np.array([], dtype=np.int8)


class ParquetStreamReader:

    def __init__(self, parquet_files, columns, read_batch_rows=100_000):
        self.parquet_files = parquet_files
        self.columns = columns
        self.read_batch_rows = read_batch_rows
        self._file_iter = iter(parquet_files)
        self._batch_iter = None
        self._buffer = None  # leftover ndarray from a previous pull

    def _next_source_batch(self):
        while True:
            if self._batch_iter is None:
                try:
                    path = next(self._file_iter)
                except StopIteration:
                    return None
                dataset = ds.dataset(path, format="parquet")
                scanner = dataset.scanner(columns=self.columns, batch_size=self.read_batch_rows)
                self._batch_iter = scanner.to_batches()
            try:
                return next(self._batch_iter)
            except StopIteration:
                self._batch_iter = None  # exhausted this file, move to the next

    def next_rows(self, n):
        chunks = []
        have = 0
        if self._buffer is not None:
            chunks.append(self._buffer)
            have = self._buffer.shape[0]
            self._buffer = None

        while have < n:
            batch = self._next_source_batch()
            if batch is None:
                break
            arr = batch.to_pandas().to_numpy(dtype=np.float32)
            chunks.append(arr)
            have += arr.shape[0]

        if not chunks:
            return np.empty((0, len(self.columns)), dtype=np.float32)

        combined = np.concatenate(chunks, axis=0) if len(chunks) > 1 else chunks[0]
        result = combined[:n]
        leftover = combined[n:]
        self._buffer = leftover if leftover.shape[0] > 0 else None
        return result


def files_for_range(file_boundaries, start, end):
    files = []
    for name, f_start, f_end in file_boundaries:
        if f_start < end and f_end > start:
            files.append(name)
    return files


def to_numpy(arr):
    result = arr.compute() if hasattr(arr, "compute") else arr
    return result


def _ensure_dask(X, n_chunks=1):
    if isinstance(X, da.Array):
        return X
    rows = max(X.shape[0] // n_chunks, 1)
    return da.from_array(X, chunks=(rows, X.shape[1]))


def distance_to_centroids_numpy(block, centers=None):
    diffs = block[:, np.newaxis, :] - centers[np.newaxis, :, :]
    distances = np.linalg.norm(diffs, axis=2).min(axis=1)
    return distances


def kmeans_distance_score(model, X, backend="numpy"):
    centers = model.cluster_centers_
    if hasattr(centers, "compute"):
        centers = centers.compute()
    X_np = to_numpy(X)
    return distance_to_centroids_numpy(X_np, centers=centers)


class ModelState:
    def __init__(self, algorithm, backend, n_clusters, random_state, min_cluster_size, n_estimators, distance_backend, sdo_chunksize, sdo_k=None, n_workers=1):
        self.algorithm = algorithm
        self.backend = backend
        self.n_clusters = n_clusters
        self.random_state = random_state
        self.min_cluster_size = min_cluster_size
        self.n_estimators = n_estimators
        self.distance_backend = distance_backend
        self.sdo_chunksize = sdo_chunksize
        self.sdo_k = sdo_k
        self.n_workers = n_workers
        self.model = None


def build_model(config, n_workers=1):
    state = ModelState(
        algorithm=config["algorithm"],
        backend=config.get("backend", "numpy"),
        n_clusters=config.get("k"),
        random_state=config.get("rseed"),
        min_cluster_size=config.get("min_cluster_size", 10),
        n_estimators=config.get("n_estimators", 200),
        distance_backend=config.get("distance_backend", "numpy"),
        sdo_chunksize=config.get("sdo_chunksize", 10000),
        sdo_k=config.get("sdo_k", None),
        n_workers=n_workers
    )
    return state


def initialize(state, X):
    if state.algorithm == "SDOclust":
        if state.sdo_k is None:
            state.model = sdo.SDOclust(chi=3, e=1, chunksize=state.sdo_chunksize, backend=state.backend, n_jobs=-1)
        else:
            state.model = sdo.SDOclust(chi=3, e=1, k=state.sdo_k, chunksize=state.sdo_chunksize, backend=state.backend, n_jobs=-1)
        output = state.model.fit_predict(X, return_membership=False)

    elif state.algorithm in ("SDO", "SDO-frozen"):
        if state.sdo_k is None:
            state.model = sdo.SDO(chunksize=state.sdo_chunksize, backend=state.backend, n_jobs=-1)
        else:
            state.model = sdo.SDO(k=state.sdo_k, chunksize=state.sdo_chunksize, backend=state.backend, n_jobs=-1)
        output = state.model.fit_predict(X)

    elif state.algorithm == "Wp_DaskKMeans":
        Xd = _ensure_dask(X, state.n_workers)
        state.model = DaskKMeans(n_clusters=state.n_clusters, random_state=state.random_state)
        state.model.fit(Xd)
        output = state.model.predict(Xd)

    elif state.algorithm in ("Wp_DaskKMeansAD", "Wp_DaskKMeansAD-frozen"):
        Xd = _ensure_dask(X, state.n_workers)
        state.model = DaskKMeans(n_clusters=state.n_clusters, random_state=state.random_state)
        state.model.fit(Xd)
        output = kmeans_distance_score(state.model, X, backend=state.distance_backend)

    elif state.algorithm == "MiniBatchKMeans":
        X = to_numpy(X)
        state.model = MiniBatchKMeans(n_clusters=state.n_clusters, random_state=state.random_state)
        state.model.fit(X)
        output = state.model.predict(X)

    elif state.algorithm == "fast-HDBSCAN":
        X = to_numpy(X)
        state.model = HDBSCAN(min_cluster_size=state.min_cluster_size)
        state.model.fit(X)
        output = state.model.labels_

    elif state.algorithm in ("IsolationForest", "IsolationForest-frozen"):
        X = to_numpy(X)
        state.model = IsolationForest(n_estimators=state.n_estimators, random_state=state.random_state, n_jobs=-1)
        state.model.fit(X)
        output = -state.model.score_samples(X)

    else:
        raise ValueError(f"Unknown algorithm: {state.algorithm}")

    return output


def update(state, X):
    if state.algorithm == "SDO":
        output = state.model.fit_predict(X)

    elif state.algorithm == "SDO-frozen":
        output = state.model.predict(X)

    elif state.algorithm == "SDOclust":
        output = state.model.update_predict(X, return_membership=False)

    elif state.algorithm == "Wp_DaskKMeans":
        Xd = _ensure_dask(X, state.n_workers)
        centers = state.model.cluster_centers_
        state.model = DaskKMeans(n_clusters=state.n_clusters, init=centers, n_init=1)
        state.model.fit(Xd)
        output = state.model.predict(Xd)

    elif state.algorithm == "Wp_DaskKMeansAD":
        Xd = _ensure_dask(X, state.n_workers)
        centers = state.model.cluster_centers_
        state.model = DaskKMeans(n_clusters=state.n_clusters, init=centers, n_init=1)
        state.model.fit(Xd)
        output = kmeans_distance_score(state.model, X, backend=state.distance_backend)

    elif state.algorithm == "Wp_DaskKMeansAD-frozen":
        output = kmeans_distance_score(state.model, X, backend=state.distance_backend)

    elif state.algorithm == "MiniBatchKMeans":
        X = to_numpy(X)
        state.model.partial_fit(X)
        output = state.model.predict(X)

    elif state.algorithm == "fast-HDBSCAN":
        X = to_numpy(X)
        state.model = HDBSCAN(min_cluster_size=state.min_cluster_size)
        state.model.fit(X)
        output = state.model.labels_

    elif state.algorithm == "IsolationForest":
        X = to_numpy(X)
        state.model = IsolationForest(n_estimators=state.n_estimators, random_state=state.random_state, n_jobs=-1)
        state.model.fit(X)
        output = -state.model.score_samples(X)

    elif state.algorithm == "IsolationForest-frozen":
        X = to_numpy(X)
        output = -state.model.score_samples(X)

    else:
        raise ValueError(f"Unknown algorithm: {state.algorithm}")

    return output


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_usage.append(psutil.cpu_percent(interval=None))
                time.sleep(0.05)

        t_monitor = threading.Thread(target=monitor)
        t_monitor.start()
        t0 = time.perf_counter()

        output = func(*args, **kwargs)

        t_total = time.perf_counter() - t0
        running = False
        t_monitor.join()

        metrics = {
            "time_sec": t_total,
            "peak_mem_mb": peak_rss / 1024**2,
            "cpu_avg_percent": float(np.mean(cpu_usage)),
            "cpu_max_percent": float(np.max(cpu_usage)),
            "num_threads": threading.active_count(),
            "num_process_threads": process.num_threads()
        }
        result = (output, metrics)
        return result
    return wrapper


@monitor_resources
def controlled_process(state, X, task="initialize"):
    if task == "initialize":
        output = initialize(state, X)
    else:
        output = update(state, X)
    return output


def compute_metrics(y_true, y_pred, task, metrics_list, outlier_frac=None):
    scores = {}
    if task == "clustering":
        mask = y_true >= 0
        if "ARI" in metrics_list:
            scores["ARI"] = adjusted_rand_score(y_true[mask], y_pred[mask])
        if "AMI" in metrics_list:
            scores["AMI"] = adjusted_mutual_info_score(y_true[mask], y_pred[mask])
        return scores

    n_pos = int(np.sum(y_true == 1))
    n_neg = len(y_true) - n_pos
    if n_pos == 0 or n_neg == 0:
        for m in metrics_list:
            scores[m] = np.nan
        return scores

    bin_preds = None
    if "mcc" in metrics_list or "ami" in metrics_list:
        threshold = np.quantile(y_pred, 1 - outlier_frac)
        bin_preds = (y_pred >= threshold).astype(int)
    scores = fm.get_indices(y_true, y_pred, bin_preds=bin_preds, metrics=set(metrics_list))
    return scores


def run_stream(parquet_files, features, y_true, state, task, initial_chunk, batch_size, metrics_list,
               file_boundaries, outlier_frac=None, eval_window_size=None, read_batch_rows=100_000):
    results = []
    reader = ParquetStreamReader(parquet_files, features, read_batch_rows=min(read_batch_rows, batch_size))

    window_y = deque(maxlen=eval_window_size) if eval_window_size else None
    window_pred = deque(maxlen=eval_window_size) if eval_window_size else None

    n_samples = len(y_true)

    # --- initial chunk ---
    X_init = reader.next_rows(initial_chunk)
    n_init = X_init.shape[0]
    y_init = y_true[:n_init]

    y_pred_init, run_metrics = controlled_process(state, X_init, task="initialize")
    y_pred_init = to_numpy(y_pred_init)

    batch_scores = compute_metrics(y_init, y_pred_init, task, metrics_list, outlier_frac)
    scores = {f"batch_{k}": v for k, v in batch_scores.items()}

    if task != "clustering":
        scores["batch_pct_pos"] = 100 * np.mean(y_init == 1)

    if window_y is not None:
        window_y.extend(y_init)
        window_pred.extend(y_pred_init)
        y_win = np.array(window_y)
        pred_win = np.array(window_pred)
        scores.update(compute_metrics(y_win, pred_win, task, metrics_list, outlier_frac))
        if task != "clustering":
            scores["window_pct_pos"] = 100 * np.mean(y_win == 1)
    else:
        scores.update(batch_scores)

    init_files = files_for_range(file_boundaries, 0, n_init)
    log = {"phase": "init", "start": 0, "end": n_init, "files": ";".join(init_files)} | scores | run_metrics
    results.append(log)
    print(f"Init chunk (0-{n_init}) [{log['files']}] -> {scores}")

    del X_init, y_init, y_pred_init
    gc.collect()

    # --- streaming batches ---
    start = n_init
    while start < n_samples:
        end = min(start + batch_size, n_samples)
        X_batch = reader.next_rows(end - start)
        if X_batch.shape[0] == 0:
            break  # stream exhausted early (shouldn't happen if row counts line up)
        end = start + X_batch.shape[0]  # in case fewer rows were returned than expected
        y_batch = y_true[start:end]

        y_pred_batch, run_metrics = controlled_process(state, X_batch, task="update")
        y_pred_batch = to_numpy(y_pred_batch)

        batch_scores = compute_metrics(y_batch, y_pred_batch, task, metrics_list, outlier_frac)
        scores = {f"batch_{k}": v for k, v in batch_scores.items()}
        if task != "clustering":
            scores["batch_pct_pos"] = 100 * np.mean(y_batch == 1)

        if window_y is not None:
            window_y.extend(y_batch)
            window_pred.extend(y_pred_batch)
            y_win = np.array(window_y)
            pred_win = np.array(window_pred)
            scores.update(compute_metrics(y_win, pred_win, task, metrics_list, outlier_frac))
            if task != "clustering":
                scores["window_pct_pos"] = 100 * np.mean(y_win == 1)
        else:
            scores.update(batch_scores)

        batch_files = files_for_range(file_boundaries, start, end)
        log = {"phase": "batch", "start": start, "end": end, "files": ";".join(batch_files)} | scores | run_metrics
        results.append(log)
        print(f"Batch {start}-{end} [{log['files']}] -> {scores}")

        del X_batch, y_batch, y_pred_batch
        gc.collect()

        start = end

    result = pd.DataFrame(results)
    return result


def build_client(cluster_config):
    dask.config.set({ "distributed.worker.profile.enabled": False })
    cluster = None
    if cluster_config is None or cluster_config.get("type", "local") == "local":
        n_workers = cluster_config.get("n_workers") if cluster_config else None
        client = Client(n_workers=n_workers, local_directory="./results/dask-worker-space") if n_workers \
            else Client(local_directory="./results/dask-worker-space")
    else:
        cluster = SLURMCluster(
            queue=cluster_config["queue"],
            cores=cluster_config.get("cores", 1),
            memory=cluster_config.get("memory", "16G"),
            processes=cluster_config.get("processes", 1),
            walltime=cluster_config.get("walltime", "24:00:00"),
            job_extra_directives=cluster_config.get("job_extra_directives", []),
            local_directory="./results/dask-worker-space",  
        )
        cluster.scale(cluster_config.get("n_workers", 4))
        client = Client(cluster)  
        client.wait_for_workers(n_workers=cluster_config.get("n_workers", 4))
    return client, cluster
    

def main(experiment_path, execution_path):
    config = load_config(experiment_path)
    execution_config = load_config(execution_path)
    os.makedirs(config["output_folder"], exist_ok=True)

    dask.config.set(scheduler="distributed")
    client, cluster = build_client(execution_config.get("cluster"))
    print(client.dashboard_link)

    parquet_files = config["parquet_files"]
    features = config["features"]
    target = config["target"]

    file_boundaries, _ = get_file_boundaries(parquet_files)
    y_true = load_targets(parquet_files, target, config.get("binarize_target_0class"))

    state = build_model(config, n_workers=len(client.scheduler_info()['workers']))

    df = run_stream(
        parquet_files, features, y_true, state, config["task"],
        config["initial_chunk"], config["batch_size"], config["metrics"],
        file_boundaries, config.get("outlier_frac"), config.get("eval_window_size"),
        read_batch_rows=config.get("read_batch_rows", 100000)
    )

    df["algorithm"] = config["algorithm"]
    df["backend"] = config.get("backend", "numpy")
    df["task"] = config["task"]

    fname = f"{config['output_folder']}/results_stream_{config['task']}_{config['algorithm']}_{df['backend'].iloc[0]}.csv"
    df.to_csv(fname, index=False)
    print(f"Saved results to {fname}")

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


if __name__ == "__main__":
    if len(sys.argv) < 3:
        print("Usage: python stream_test.py <experiment.json> <execution.json>")
        sys.exit(1)
    main(sys.argv[1], sys.argv[2])
