#!/usr/bin/env python3
"""AI model training pipeline.

This script implements a practical training workflow for tabular machine learning:
- load and validate data
- split training and validation sets
- build a preprocessing + model pipeline
- evaluate on validation data
- export metrics, metadata, and the trained model artifact

Assumptions:
- Input data is a CSV file.
- A target column is present.
- Feature columns may be numeric, categorical, or mixed.
- The default task is classification.

Usage examples:
    python train_pipeline.py --data data/training.csv --target target --output-dir models
    python train_pipeline.py --data data/training.csv --target target --task regression

Notes:
- No credentials or environment-specific endpoints are required.
- This script avoids destructive actions and writes only to the chosen output directory.
"""

from __future__ import annotations

import argparse
import json
from dataclasses import dataclass, asdict
from pathlib import Path
from typing import Dict, List, Optional, Tuple

import joblib
import pandas as pd
from sklearn.compose import ColumnTransformer
from sklearn.impute import SimpleImputer
from sklearn.linear_model import LinearRegression, LogisticRegression
from sklearn.metrics import (
    accuracy_score,
    f1_score,
    mean_absolute_error,
    mean_squared_error,
    r2_score,
)
from sklearn.model_selection import train_test_split
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import OneHotEncoder, StandardScaler


@dataclass
class TrainingConfig:
    data_path: Path
    target_column: str
    output_dir: Path
    task: str = "classification"  # classification or regression
    test_size: float = 0.2
    random_state: int = 42
    drop_duplicates: bool = True


def parse_args() -> TrainingConfig:
    parser = argparse.ArgumentParser(
        description="Train and export a reproducible ML pipeline from CSV data."
    )
    parser.add_argument("--data", required=True, help="Path to input CSV file.")
    parser.add_argument(
        "--target",
        required=True,
        help="Name of the target column in the dataset.",
    )
    parser.add_argument(
        "--output-dir",
        default="models",
        help="Directory where artifacts will be written.",
    )
    parser.add_argument(
        "--task",
        choices=["classification", "regression"],
        default="classification",
        help="Training task type.",
    )
    parser.add_argument(
        "--test-size",
        type=float,
        default=0.2,
        help="Validation split fraction between 0 and 1.",
    )
    parser.add_argument(
        "--random-state",
        type=int,
        default=42,
        help="Random seed for reproducible splitting.",
    )
    parser.add_argument(
        "--keep-duplicates",
        action="store_true",
        help="Keep duplicate rows instead of dropping them.",
    )

    args = parser.parse_args()

    if not 0.0 < args.test_size < 1.0:
        parser.error("--test-size must be between 0 and 1.")

    return TrainingConfig(
        data_path=Path(args.data),
        target_column=args.target,
        output_dir=Path(args.output_dir),
        task=args.task,
        test_size=args.test_size,
        random_state=args.random_state,
        drop_duplicates=not args.keep_duplicates,
    )


def load_and_validate_data(path: Path, target_column: str, drop_duplicates: bool = True) -> pd.DataFrame:
    if not path.exists():
        raise FileNotFoundError(f"Data file not found: {path}")
    if path.suffix.lower() != ".csv":
        raise ValueError("Input data must be a CSV file.")

    df = pd.read_csv(path)
    if df.empty:
        raise ValueError("Input data is empty.")
    if target_column not in df.columns:
        raise ValueError(f"Missing target column: {target_column}")

    if drop_duplicates:
        df = df.drop_duplicates().reset_index(drop=True)

    if df[target_column].isna().any():
        raise ValueError("Target column contains missing values.")

    feature_columns = [c for c in df.columns if c != target_column]
    if not feature_columns:
        raise ValueError("No feature columns found after excluding the target.")

    # Basic sanity checks for training readiness.
    if df.isna().all(axis=1).any():
        raise ValueError("At least one row contains only missing values.")

    return df


def split_data(
    df: pd.DataFrame,
    target_column: str,
    task: str,
    test_size: float,
    random_state: int,
) -> Tuple[pd.DataFrame, pd.DataFrame, pd.Series, pd.Series]:
    X = df.drop(columns=[target_column])
    y = df[target_column]

    stratify = y if task == "classification" and y.nunique() > 1 else None

    return train_test_split(
        X,
        y,
        test_size=test_size,
        random_state=random_state,
        stratify=stratify,
    )


def build_pipeline(X: pd.DataFrame, task: str) -> Pipeline:
    numeric_features = X.select_dtypes(include=["number", "bool"]).columns.tolist()
    categorical_features = [c for c in X.columns if c not in numeric_features]

    numeric_transformer = Pipeline(
        steps=[
            ("imputer", SimpleImputer(strategy="median")),
            ("scaler", StandardScaler()),
        ]
    )

    categorical_transformer = Pipeline(
        steps=[
            ("imputer", SimpleImputer(strategy="most_frequent")),
            ("encoder", OneHotEncoder(handle_unknown="ignore")),
        ]
    )

    transformers = []
    if numeric_features:
        transformers.append(("num", numeric_transformer, numeric_features))
    if categorical_features:
        transformers.append(("cat", categorical_transformer, categorical_features))

    if not transformers:
        raise ValueError("No usable feature columns found for preprocessing.")

    preprocessor = ColumnTransformer(transformers=transformers)

    if task == "classification":
        estimator = LogisticRegression(max_iter=1000)
    else:
        estimator = LinearRegression()

    return Pipeline(
        steps=[
            ("preprocessor", preprocessor),
            ("model", estimator),
        ]
    )


def evaluate_model(model: Pipeline, X_val: pd.DataFrame, y_val: pd.Series, task: str) -> Dict[str, float]:
    predictions = model.predict(X_val)

    if task == "classification":
        metrics = {
            "accuracy": accuracy_score(y_val, predictions),
            "f1_weighted": f1_score(y_val, predictions, average="weighted"),
        }
    else:
        mse = mean_squared_error(y_val, predictions)
        metrics = {
            "mae": mean_absolute_error(y_val, predictions),
            "rmse": mse ** 0.5,
            "r2": r2_score(y_val, predictions),
        }

    return metrics


def write_artifacts(
    output_dir: Path,
    model: Pipeline,
    metrics: Dict[str, float],
    config: TrainingConfig,
    feature_columns: List[str],
    row_count: int,
) -> None:
    output_dir.mkdir(parents=True, exist_ok=True)

    model_path = output_dir / "trained-model.joblib"
    metrics_path = output_dir / "metrics.json"
    metadata_path = output_dir / "metadata.json"

    joblib.dump(model, model_path)

    with metrics_path.open("w", encoding="utf-8") as f:
        json.dump(metrics, f, indent=2, sort_keys=True)

    metadata = {
        "task": config.task,
        "target_column": config.target_column,
        "feature_columns": feature_columns,
        "input_rows": row_count,
        "test_size": config.test_size,
        "random_state": config.random_state,
    }
    with metadata_path.open("w", encoding="utf-8") as f:
        json.dump(metadata, f, indent=2, sort_keys=True)

    print(f"Saved model to: {model_path}")
    print(f"Saved metrics to: {metrics_path}")
    print(f"Saved metadata to: {metadata_path}")


def main() -> int:
    config = parse_args()

    df = load_and_validate_data(
        path=config.data_path,
        target_column=config.target_column,
        drop_duplicates=config.drop_duplicates,
    )

    X_train, X_val, y_train, y_val = split_data(
        df=df,
        target_column=config.target_column,
        task=config.task,
        test_size=config.test_size,
        random_state=config.random_state,
    )

    pipeline = build_pipeline(X_train, config.task)
    pipeline.fit(X_train, y_train)

    metrics = evaluate_model(pipeline, X_val, y_val, config.task)
    print("Validation metrics:")
    print(json.dumps(metrics, indent=2, sort_keys=True))

    feature_columns = [c for c in df.columns if c != config.target_column]
    write_artifacts(
        output_dir=config.output_dir,
        model=pipeline,
        metrics=metrics,
        config=config,
        feature_columns=feature_columns,
        row_count=len(df),
    )

    return 0


if __name__ == "__main__":
    raise SystemExit(main())