from __future__ import annotations

from datetime import UTC, datetime
from pathlib import Path

import pandas as pd

from psx_signal.config import Settings
from psx_signal.evaluation import classification_metrics
from psx_signal.models import create_model
from psx_signal.pipelines.common import feature_columns, frame_fingerprint, write_json


def train_model(
    dataset: pd.DataFrame,
    settings: Settings,
    model_name: str,
    horizon: int,
    artifact_path: str | Path,
) -> dict[str, object]:
    label_column = f"label_{horizon}d"
    clean = dataset.dropna(subset=[label_column]).sort_values("date").copy()
    columns = feature_columns(clean)
    if not columns:
        raise ValueError("No numeric feature columns found")
    sessions = pd.DatetimeIndex(clean["date"].unique()).sort_values()
    validation_size = settings.walk_forward.validation_sessions
    if len(sessions) <= validation_size:
        raise ValueError(f"Need more than {validation_size} sessions for train/validation")
    split_date = sessions[-validation_size]
    train = clean[clean["date"] < split_date]
    validation = clean[clean["date"] >= split_date]
    model = create_model(model_name, settings).fit(train[columns], train[label_column])
    probabilities = model.predict_proba(validation[columns])
    metrics = classification_metrics(
        validation[label_column].astype(int).to_numpy(), model.predict(validation[columns]), probabilities
    )
    final_model = create_model(model_name, settings).fit(clean[columns], clean[label_column])
    final_model.save(artifact_path)
    model_version = f"{model_name}-{horizon}d-{datetime.now(UTC).strftime('%Y%m%dT%H%M%SZ')}"
    metadata: dict[str, object] = {
        "model_version": model_version,
        "model_name": model_name,
        "horizon": horizon,
        "feature_version": "m1-v1",
        "feature_columns": columns,
        "dataset_hash": frame_fingerprint(clean, ["symbol", "date", label_column] + columns),
        "config_hash": settings.fingerprint(),
        "training_date_range": [str(train["date"].min().date()), str(train["date"].max().date())],
        "validation_date_range": [str(validation["date"].min().date()), str(validation["date"].max().date())],
        "hyperparameters": settings.models.get(model_name, {}),
        "validation_metrics": metrics,
    }
    providers = sorted(clean["provider"].dropna().astype(str).unique()) if "provider" in clean else []
    trusts = sorted(clean["provider_trust"].dropna().astype(str).unique()) if "provider_trust" in clean else []
    metadata["training_provider"] = providers[0] if len(providers) == 1 else providers
    metadata["training_provider_trust"] = trusts[0] if len(trusts) == 1 else trusts
    if hasattr(final_model, "feature_importance"):
        metadata["feature_importance"] = final_model.feature_importance()
    write_json(metadata, f"{artifact_path}.json")
    return metadata
