import numpy as np
import pandas as pd

from models.lgbm_model import LGBMModel
from models.xgb_model import XGBModel


class EnsembleModel:
    def __init__(self, weights: dict = None):
        self.lgbm = LGBMModel()
        self.xgb = XGBModel()
        self.weights = weights or {"lgbm": 0.6, "xgb": 0.4}

    def train(self, X: pd.DataFrame, y: pd.Series, eval_pct: float = 0.2):
        lgbm_metrics = self.lgbm.train(X, y, eval_pct)
        xgb_metrics = self.xgb.train(X, y, eval_pct)
        return {"lgbm": lgbm_metrics, "xgb": xgb_metrics}

    def predict_proba(self, X: pd.DataFrame) -> np.ndarray:
        lgbm_proba = self.lgbm.predict_proba(X)
        xgb_proba = self.xgb.predict_proba(X)
        w_lgbm = self.weights["lgbm"]
        w_xgb = self.weights["xgb"]
        return lgbm_proba * w_lgbm + xgb_proba * w_xgb

    def predict(self, X: pd.DataFrame, threshold: float = 0.5) -> np.ndarray:
        proba = self.predict_proba(X)
        return (proba >= threshold).astype(int)

    def score_stocks(self, X: pd.DataFrame) -> pd.DataFrame:
        proba = self.predict_proba(X)
        result = pd.DataFrame({
            "ml_score": proba,
            "ml_rank": pd.Series(proba).rank(ascending=False).astype(int),
        }, index=X.index)
        result["ml_score_norm"] = (result["ml_score"] - result["ml_score"].min()) / (
            result["ml_score"].max() - result["ml_score"].min() + 1e-10
        )
        return result

    def feature_importance(self) -> pd.DataFrame:
        lgbm_imp = self.lgbm.feature_importance()
        xgb_imp = self.xgb.feature_importance()
        if lgbm_imp.empty:
            return xgb_imp
        if xgb_imp.empty:
            return lgbm_imp

        merged = lgbm_imp.merge(xgb_imp, on="feature", suffixes=("_lgbm", "_xgb"))
        merged["importance"] = (
            merged["importance_lgbm"] * self.weights["lgbm"]
            + merged["importance_xgb"] * self.weights["xgb"]
        )
        return merged[["feature", "importance"]].sort_values("importance", ascending=False)

    def save(self):
        self.lgbm.save("ensemble_lgbm")
        self.xgb.save("ensemble_xgb")

    def load(self):
        self.lgbm.load("ensemble_lgbm")
        self.xgb.load("ensemble_xgb")
