import json
from pathlib import Path

import lightgbm as lgb
import numpy as np
import pandas as pd
from sklearn.model_selection import TimeSeriesSplit

from config import BASE_DIR

MODEL_DIR = BASE_DIR / "saved_models"
MODEL_DIR.mkdir(exist_ok=True)


class LGBMModel:
    def __init__(self, params: dict = None):
        self.params = params or {
            "objective": "binary",
            "metric": "auc",
            "boosting_type": "gbdt",
            "num_leaves": 31,
            "learning_rate": 0.05,
            "feature_fraction": 0.8,
            "bagging_fraction": 0.8,
            "bagging_freq": 5,
            "verbose": -1,
            "n_estimators": 300,
            "early_stopping_rounds": 30,
        }
        self.model = None
        self.feature_names = None

    def train(self, X: pd.DataFrame, y: pd.Series, eval_pct: float = 0.2):
        self.feature_names = X.columns.tolist()
        split = int(len(X) * (1 - eval_pct))
        X_train, X_val = X.iloc[:split], X.iloc[split:]
        y_train, y_val = y.iloc[:split], y.iloc[split:]

        n_est = self.params.pop("n_estimators", 300)
        early_stop = self.params.pop("early_stopping_rounds", 30)

        self.model = lgb.LGBMClassifier(n_estimators=n_est, **self.params)
        self.model.fit(
            X_train, y_train,
            eval_X=X_val, eval_y=y_val,
            callbacks=[lgb.early_stopping(early_stop), lgb.log_evaluation(0)],
        )
        return self._eval(X_val, y_val)

    def predict_proba(self, X: pd.DataFrame) -> np.ndarray:
        if self.model is None:
            raise ValueError("Model not trained")
        if hasattr(self, "_booster"):
            return self._booster.predict(X)
        return self.model.predict_proba(X)[:, 1]

    def predict(self, X: pd.DataFrame) -> np.ndarray:
        if self.model is None:
            raise ValueError("Model not trained")
        if hasattr(self, "_booster"):
            proba = self._booster.predict(X)
            return (proba >= 0.5).astype(int)
        return self.model.predict(X)

    def feature_importance(self) -> pd.DataFrame:
        if self.model is None:
            return pd.DataFrame()
        if hasattr(self, "_booster"):
            imp = self._booster.feature_importance(importance_type="gain")
            return pd.DataFrame({
                "feature": self.feature_names,
                "importance": imp,
            }).sort_values("importance", ascending=False)
        imp = self.model.feature_importances_
        return pd.DataFrame({
            "feature": self.feature_names,
            "importance": imp,
        }).sort_values("importance", ascending=False)

    def save(self, name: str = "lgbm"):
        path = MODEL_DIR / f"{name}.txt"
        self.model.booster_.save_model(str(path))
        meta = {"feature_names": self.feature_names, "params": self.params}
        (MODEL_DIR / f"{name}_meta.json").write_text(json.dumps(meta))

    def load(self, name: str = "lgbm"):
        path = MODEL_DIR / f"{name}.txt"
        if not path.exists():
            return False
        self._booster = lgb.Booster(model_file=str(path))
        self.model = True
        meta_path = MODEL_DIR / f"{name}_meta.json"
        if meta_path.exists():
            meta = json.loads(meta_path.read_text())
            self.feature_names = meta.get("feature_names")
        return True

    def _eval(self, X_val, y_val):
        from sklearn.metrics import roc_auc_score, accuracy_score, classification_report
        proba = self.predict_proba(X_val)
        pred = (proba >= 0.5).astype(int)
        return {
            "auc": roc_auc_score(y_val, proba),
            "accuracy": accuracy_score(y_val, pred),
        }


class WalkForwardTrainer:
    def __init__(self, model_cls=LGBMModel, n_splits=5, gap=5):
        self.model_cls = model_cls
        self.n_splits = n_splits
        self.gap = gap

    def train(self, X: pd.DataFrame, y: pd.Series):
        tscv = TimeSeriesSplit(n_splits=self.n_splits)
        results = []
        models = []

        for fold, (train_idx, val_idx) in enumerate(tscv.split(X)):
            X_train = X.iloc[train_idx]
            y_train = y.iloc[train_idx]
            X_val = X.iloc[val_idx[self.gap:]]
            y_val = y.iloc[val_idx[self.gap:]]

            model = self.model_cls()
            split = int(len(X_train) * 0.9)
            model.model = None
            model.feature_names = X.columns.tolist()

            n_est = model.params.pop("n_estimators", 300)
            early_stop = model.params.pop("early_stopping_rounds", 30)

            model.model = lgb.LGBMClassifier(n_estimators=n_est, **model.params)
            model.model.fit(
                X_train.iloc[:split], y_train.iloc[:split],
                eval_X=X_train.iloc[split:], eval_y=y_train.iloc[split:],
                callbacks=[lgb.early_stopping(early_stop), lgb.log_evaluation(0)],
            )

            eval_r = model._eval(X_val, y_val)
            eval_r["fold"] = fold
            eval_r["train_size"] = len(X_train)
            eval_r["val_size"] = len(X_val)
            results.append(eval_r)
            models.append(model)

            model.params["n_estimators"] = n_est
            model.params["early_stopping_rounds"] = early_stop

        return pd.DataFrame(results), models
