```python
#!/usr/bin/env python3
"""ml_starter_workflow.py

A practical starter workflow for learning machine learning safely and reproducibly.

What this script does:
- Loads a CSV dataset
- Validates that the target column exists
- Prints basic data quality checks
- Builds a baseline model
- Evaluates the model with an appropriate split strategy
- Reports metrics and common failure signals

Supported tasks:
- classification
- regression

This script is intentionally vendor-neutral and safe by default.
It makes no external network calls and performs no destructive actions.

Examples:
  python ml_starter_workflow.py --data data.csv --target target --task classification
  python ml_starter_workflow.py --data data.csv --target price --task regression --test-size 0.2
"""

from __future__ import annotations

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

import numpy as np
import pandas as pd
from sklearn.compose import ColumnTransformer
from sklearn.dummy import DummyClassifier, DummyRegressor
from sklearn.ensemble import RandomForestClassifier, RandomForestRegressor
from sklearn.impute import SimpleImputer
from sklearn.linear_model import LogisticRegression, LinearRegression
from sklearn.metrics import (
    accuracy_score,
    f1_score,
    mean_absolute_error,
    mean_squared_error,
    precision_score,
    r2_score,
    recall_score,
    roc_auc_score,
)
from sklearn.model_selection import GroupShuffleSplit, StratifiedShuffleSplit, train_test_split
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import OneHotEncoder, StandardScaler


@dataclass
class SplitResult:
    X_train: pd.DataFrame
    X_val: pd.DataFrame
    X_test: pd.DataFrame
    y_train: pd.Series
    y_val: pd.Series
    y_test: pd.Series


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Inspect a dataset and run a baseline machine learning workflow."
    )
    parser.add_argument("--data", required=True, help="Path to input CSV file")
    parser.add_argument("--target", required=True, help="Target column name")
    parser.add_argument(
        "--task",
        required=True,
        choices=["classification", "regression"],
        help="Type of machine learning task",
    )
    parser.add_argument(
        "--group-column",
        default=None,
        help="Optional column for grouped splitting to avoid entity leakage",
    )
    parser.add_argument(
        "--time-column",
        default=None,
        help="Optional column for time-based sorting before splitting",
    )
    parser.add_argument(
        "--test-size",
        type=float,
        default=0.2,
        help="Fraction of data reserved for test set (default: 0.2)",
    )
    parser.add_argument(
        "--val-size",
        type=float,
        default=0.2,
        help="Fraction of remaining training data reserved for validation set (default: 0.2)",
    )
    parser.add_argument(
        "--random-state",
        type=int,
        default=42,
        help="Random seed for reproducibility",
    )
    parser.add_argument(
        "--output-metrics",
        default=None,
        help="Optional path to write metrics as JSON",
    )
    return parser.parse_args()


def validate_args(args: argparse.Namespace) -> None:
    data_path = Path(args.data)
    if not data_path.exists():
        raise FileNotFoundError(f"Data file not found: {data_path}")
    if data_path.suffix.lower() != ".csv":
        raise ValueError("Only CSV input is supported by this starter workflow.")
    if not (0.0 < args.test_size < 1.0):
        raise ValueError("--test-size must be between 0 and 1.")
    if not (0.0 < args.val_size < 1.0):
        raise ValueError("--val-size must be between 0 and 1.")
    if args.test_size + args.val_size >= 1.0:
        raise ValueError("--test-size plus --val-size must be less than 1.")
    if args.group_column and args.time_column:
        raise ValueError("Use either --group-column or --time-column, not both.")


def load_data(path: str) -> pd.DataFrame:
    return pd.read_csv(path)


def print_data_summary(df: pd.DataFrame, target: str) -> None:
    print("\n=== Data Summary ===")
    print("rows:", len(df))
    print("columns:", len(df.columns))

    missing = df.isna().sum().sort_values(ascending=False)
    print("\nTop missing values:")
    print(missing.head(10).to_string())

    print("\nduplicate rows:", int(df.duplicated().sum()))

    if target in df.columns:
        print("\ntarget distribution:")
        print(df[target].value_counts(dropna=False).to_string())
    else:
        print(f"\nTarget column '{target}' not found in data summary.")


def split_features_target(df: pd.DataFrame, target: str) -> Tuple[pd.DataFrame, pd.Series]:
    if target not in df.columns:
        raise KeyError(f"Target column '{target}' does not exist in the dataset.")
    y = df[target]
    X = df.drop(columns=[target])
    return X, y


def safe_time_sort(df: pd.DataFrame, time_column: str) -> pd.DataFrame:
    if time_column not in df.columns:
        raise KeyError(f"Time column '{time_column}' does not exist in the dataset.")
    sorted_df = df.sort_values(by=time_column, kind="mergesort").reset_index(drop=True)
    return sorted_df


def make_initial_split(
    X: pd.DataFrame,
    y: pd.Series,
    task: str,
    test_size: float,
    val_size: float,
    random_state: int,
    group_column: Optional[str] = None,
    time_column: Optional[str] = None,
) -> SplitResult:
    if time_column is not None:
        ordered = X.copy()
        ordered["__target__"] = y.values
        ordered = safe_time_sort(ordered, time_column)
        y_sorted = ordered.pop("__target__")
        X_sorted = ordered

        n = len(X_sorted)
        test_n = max(1, int(round(n * test_size)))
        val_n = max(1, int(round((n - test_n) * val_size)))
        train_end = n - test_n - val_n
        if train_end <= 0:
            raise ValueError("Dataset is too small for the requested split sizes.")

        X_train = X_sorted.iloc[:train_end].copy()
        y_train = y_sorted.iloc[:train_end].copy()
        X_val = X_sorted.iloc[train_end:train_end + val_n].copy()
        y_val = y_sorted.iloc[train_end:train_end + val_n].copy()
        X_test = X_sorted.iloc[train_end + val_n:].copy()
        y_test = y_sorted.iloc[train_end + val_n:].copy()
        return SplitResult(X_train, X_val, X_test, y_train, y_val, y_test)

    if group_column is not None:
        if group_column not in X.columns:
            raise KeyError(f"Group column '{group_column}' does not exist in features.")
        groups = X[group_column]
        X_wo_group = X.drop(columns=[group_column])

        splitter = GroupShuffleSplit(n_splits=1, test_size=test_size, random_state=random_state)
        train_idx, test_idx = next(splitter.split(X_wo_group, y, groups=groups))
        X_train_full, X_test = X_wo_group.iloc[train_idx].copy(), X_wo_group.iloc[test_idx].copy()
        y_train_full, y_test = y.iloc[train_idx].copy(), y.iloc[test_idx].copy()

        if task == "classification":
            strat = StratifiedShuffleSplit(n_splits=1, test_size=val_size, random_state=random_state)
            train_idx2, val_idx = next(strat.split(X_train_full, y_train_full))
            X_train = X_train_full.iloc[train_idx2].copy()
            y_train = y_train_full.iloc[train_idx2].copy()
            X_val = X_train_full.iloc[val_idx].copy()
            y_val = y_train_full.iloc[val_idx].copy()
        else:
            n = len(X_train_full)
            val_n = max(1, int(round(n * val_size)))
            X_train = X_train_full.iloc[:-val_n].copy()
            y_train = y_train_full.iloc[:-val_n].copy()
            X_val = X_train_full.iloc[-val_n:].copy()
            y_val = y_train_full.iloc[-val_n:].copy()

        return SplitResult(X_train, X_val, X_test, y_train, y_val, y_test)

    stratify = y if task == "classification" else None
    X_train_full, X_test, y_train_full, y_test = train_test_split(
        X,
        y,
        test_size=test_size,
        random_state=random_state,
        stratify=stratify,
    )
    stratify2 = y_train_full if task == "classification" else None
    val_fraction_of_train = val_size
    X_train, X_val, y_train, y_val = train_test_split(
        X_train_full,
        y_train_full,
        test_size=val_fraction_of_train,
        random_state=random_state,
        stratify=stratify2,
    )
    return SplitResult(X_train, X_val, X_test, y_train, y_val, y_test)


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

    numeric_pipeline = Pipeline(
        steps=[
            ("imputer", SimpleImputer(strategy="median")),
            ("scaler", StandardScaler()),
        ]
    )
    categorical_pipeline = Pipeline(
        steps=[
            ("imputer", SimpleImputer(strategy="most_frequent")),
            ("onehot", OneHotEncoder(handle_unknown="ignore")),
        ]
    )

    return ColumnTransformer(
        transformers=[
            ("num", numeric_pipeline, numeric_features),
            ("cat", categorical_pipeline, categorical_features),
        ],
        remainder="drop",
    )


def choose_models(task: str):
    if task == "classification":
        return {
            "naive": DummyClassifier(strategy="most_frequent"),
            "baseline": LogisticRegression(max_iter=1000),
            "tree": RandomForestClassifier(n_estimators=200, random_state=42),
        }
    return {
        "naive": DummyRegressor(strategy="mean"),
        "baseline": LinearRegression(),
        "tree": RandomForestRegressor(n_estimators=200, random_state=42),
    }


def evaluate_classification(y_true: pd.Series, y_pred: np.ndarray, y_proba: Optional[np.ndarray] = None) -> Dict[str, float]:
    metrics = {
        "accuracy": float(accuracy_score(y_true, y_pred)),
        "precision": float(precision_score(y_true, y_pred, average="weighted", zero_division=0)),
        "recall": float(recall_score(y_true, y_pred, average="weighted", zero_division=0)),
        "f1": float(f1_score(y_true, y_pred, average="weighted", zero_division=0)),
    }
    if y_proba is not None:
        try:
            if y_proba.ndim == 2 and y_proba.shape[1] > 1:
                metrics["roc_auc"] = float(roc_auc_score(y_true, y_proba[:, 1]))
        except Exception:
            pass
    return metrics


def evaluate_regression(y_true: pd.Series, y_pred: np.ndarray) -> Dict[str, float]:
    return {
        "mae": float(mean_absolute_error(y_true, y_pred)),
        "rmse": float(mean_squared_error(y_true, y_pred, squared=False)),
        "r2": float(r2_score(y_true, y_pred)),
    }


def fit_and_score_model(task: str, model, X_train, y_train, X_val, y_val) -> Dict[str, float]:
    preprocessor = build_preprocessor(X_train)
    pipe = Pipeline(steps=[("preprocess", preprocessor), ("model", model)])
    pipe.fit(X_train, y_train)

    y_pred = pipe.predict(X_val)
    y_proba = None
    if task == "classification" and hasattr(pipe.named_steps["model"], "predict_proba"):
        try:
            y_proba = pipe.predict_proba(X_val)
        except Exception:
            y_proba = None

    if task == "classification":
        return evaluate_classification(y_val, y_pred, y_proba)
    return evaluate_regression(y_val, y_pred)


def main() -> int:
    args = parse_args()
    validate_args(args)

    df = load_data(args.data)
    print_data_summary(df, args.target)

    X, y = split_features_target(df, args.target)
    splits = make_initial_split(
        X=X,
        y=y,
        task=args.task,
        test_size=args.test_size,
        val_size=args.val_size,
        random_state=args.random_state,
        group_column=args.group_column,
        time_column=args.time_column,
    )

    print("\n=== Split Sizes ===")
    print("train:", len(splits.X_train))
    print("validation:", len(splits.X_val))
    print("test:", len(splits.X_test))

    models = choose_models(args.task)
    results: Dict[str, Dict[str, float]] = {}

    for name, model in models.items():
        try:
            results[name] = fit_and_score_model(
                args.task,
                model,
                splits.X_train,
                splits.y_train,
                splits.X_val,
                splits.y_val,
            )
        except Exception as exc:
            results[name] = {"error": str(exc)}

    print("\n=== Validation Results ===")
    for name, metrics in results.items():
        print(f"\n[{name}]")
        print(json.dumps(metrics, indent=2, sort_keys=True))

    if args.output_metrics:
        output_path = Path(args.output_metrics)
        output_path.parent.mkdir(parents=True, exist_ok=True)
        payload = {
            "task": args.task,
            "target": args.target,
            "random_state": args.random_state,
            "results": results,
            "split_sizes": {
                "train": len(splits.X_train),
                "validation": len(splits.X_val),
                "test": len(splits.X_test),
            },
        }
        output_path.write_text(json.dumps(payload, indent=2, sort_keys=True), encoding="utf-8")
        print(f"\nMetrics written to: {output_path}")

    print("\n=== Interpretation Checklist ===")
    print("- Does the baseline beat a trivial rule?")
    print("- Is validation performance stable and not obviously overfit?")
    print("- Are there signs of leakage, duplicate records, or impossible values?")
    print("- Would a different split strategy change the conclusion?")

    return 0


if __name__ == "__main__":
    try:
        raise SystemExit(main())
    except KeyboardInterrupt:
        print("Interrupted.", file=sys.stderr)
        raise SystemExit(130)
```