import pandas as pd
import numpy as np
from models.train import load_model, MODEL_FEATURES


def _get_features(model_artifact: dict) -> list[str]:
    """Извлекает список признаков из артефакта модели.
    Приоритет: сохранённые features > MODEL_FEATURES по умолчанию."""
    saved = model_artifact.get('features')
    if saved is not None and len(saved) > 0:
        return saved
    return MODEL_FEATURES


def predict(model_artifact: dict, df: pd.DataFrame) -> np.ndarray:
    model = model_artifact['model']
    scaler = model_artifact.get('scaler')
    features = _get_features(model_artifact)
    available = [c for c in features if c in df.columns]
    if len(available) < len(features):
        missing = set(features) - set(available)
        print(f'  [WARN] predict.py: отсутствует {len(missing)} признаков: {list(missing)[:3]}...')
    X = df[available].fillna(0)
    if scaler:
        X = scaler.transform(X)
    return model.predict(X)


def predict_proba(model_artifact: dict, df: pd.DataFrame) -> np.ndarray:
    model = model_artifact['model']
    scaler = model_artifact.get('scaler')
    features = _get_features(model_artifact)
    available = [c for c in features if c in df.columns]
    if len(available) < len(features):
        missing = set(features) - set(available)
        print(f'  [WARN] predict.py: отсутствует {len(missing)} признаков: {list(missing)[:3]}...')
    X = df[available].fillna(0)
    if scaler:
        X = scaler.transform(X)
    return model.predict_proba(X)


def generate_signal(model_artifact: dict, df: pd.DataFrame, threshold: float = 0.6) -> pd.Series:
    probs = predict_proba(model_artifact, df)
    # Защита от single-class модели (shape (N,1))
    if probs.ndim == 2 and probs.shape[1] >= 2:
        p1 = probs[:, 1]
    else:
        p1 = probs[:, 0] if probs.ndim == 2 else probs
    signals = np.where(p1 >= threshold, 1, np.where(p1 <= 1 - threshold, -1, 0))
    return pd.Series(signals, index=df.index[-len(signals):])
