"""Model and decision-support utilities for SafeEvac AI."""
from __future__ import annotations

from dataclasses import dataclass
from pathlib import Path
import os
from typing import Dict, List, Tuple

import numpy as np
import pandas as pd
from sklearn.compose import ColumnTransformer
from sklearn.ensemble import RandomForestRegressor
from sklearn.metrics import mean_absolute_error, r2_score
from sklearn.model_selection import train_test_split
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import OneHotEncoder, StandardScaler

APP_DIR = Path(__file__).resolve().parent
TRAIN_PATH = APP_DIR / "data" / "safeevac_synthetic_training.csv"
RANDOM_STATE = 42

NUMERIC_FEATURES = [
    "occupant_count", "floors", "exits_total", "blocked_exits", "stair_width_m",
    "avg_route_m", "bhv_count", "mobility_support_count",
    "sensory_cognitive_support_count", "visitor_pct", "alarm_delay_s",
    "trained_occupants_pct", "communication_channels", "night_shift", "drill_recent",
]
CATEGORICAL_FEATURES = ["sector", "hazard"]
FEATURES = NUMERIC_FEATURES + CATEGORICAL_FEATURES

PROHIBITED_INDIVIDUAL_FIELDS = {
    "sex", "gender", "ethnicity", "nationality", "religion", "income", "pclass",
    "fare", "medical_diagnosis", "disability_diagnosis", "political_belief",
}


@dataclass
class ModelMetrics:
    mae: float
    r2: float
    train_rows: int
    test_rows: int


def load_training_data() -> pd.DataFrame:
    if not TRAIN_PATH.exists():
        raise FileNotFoundError(f"Training data not found: {TRAIN_PATH}")
    return pd.read_csv(TRAIN_PATH)


def build_pipeline() -> Pipeline:
    preprocessor = ColumnTransformer([
        ("num", StandardScaler(), NUMERIC_FEATURES),
        ("cat", OneHotEncoder(handle_unknown="ignore"), CATEGORICAL_FEATURES),
    ])
    model = RandomForestRegressor(
        n_estimators=450,
        min_samples_leaf=3,
        max_features=0.8,
        random_state=RANDOM_STATE,
        n_jobs=int(os.getenv("SAFEEVAC_N_JOBS", "1")),
    )
    return Pipeline([("preprocess", preprocessor), ("model", model)])


def train_model() -> Tuple[Pipeline, ModelMetrics]:
    df = load_training_data()
    X_train, X_test, y_train, y_test = train_test_split(
        df[FEATURES], df["risk_index"], test_size=0.2, random_state=RANDOM_STATE,
        stratify=df["risk_band"],
    )
    pipe = build_pipeline()
    pipe.fit(X_train, y_train)
    pred = pipe.predict(X_test)
    metrics = ModelMetrics(
        mae=float(mean_absolute_error(y_test, pred)),
        r2=float(r2_score(y_test, pred)),
        train_rows=len(X_train), test_rows=len(X_test),
    )
    # Retrain on all available prototype rows for the interactive app.
    pipe.fit(df[FEATURES], df["risk_index"])
    return pipe, metrics


def predict_with_uncertainty(pipe: Pipeline, scenario: pd.DataFrame) -> Tuple[float, float, float]:
    """Return prediction plus tree-distribution P10/P90 as a model spread indicator."""
    transformed = pipe.named_steps["preprocess"].transform(scenario[FEATURES])
    forest = pipe.named_steps["model"]
    tree_preds = np.array([tree.predict(transformed)[0] for tree in forest.estimators_])
    pred = float(np.mean(tree_preds))
    p10, p90 = np.percentile(tree_preds, [10, 90])
    return pred, float(p10), float(p90)


def risk_band(score: float) -> str:
    if score <= 35:
        return "Low"
    if score <= 60:
        return "Medium"
    return "High"


def validate_scenario(scenario: pd.DataFrame) -> List[str]:
    issues: List[str] = []
    cols = {c.lower() for c in scenario.columns}
    forbidden = sorted(cols.intersection(PROHIBITED_INDIVIDUAL_FIELDS))
    if forbidden:
        issues.append("Prohibited individual fields detected: " + ", ".join(forbidden))
    row = scenario.iloc[0]
    if int(row["blocked_exits"]) >= int(row["exits_total"]):
        issues.append("At least one usable exit must remain in this prototype scenario.")
    if int(row["mobility_support_count"]) + int(row["sensory_cognitive_support_count"]) > int(row["occupant_count"]):
        issues.append("Support-demand counts cannot exceed total occupants.")
    return issues


def recommended_interventions(pipe: Pipeline, scenario: pd.DataFrame, base_risk: float) -> List[Dict[str, object]]:
    """Counterfactual, non-person-ranking interventions ranked by predicted risk reduction."""
    row = scenario.iloc[0].copy()
    alternatives: List[Tuple[str, str, object]] = []

    if row["blocked_exits"] > 0:
        alternatives.append(("Restore exit availability", "blocked_exits", 0))
    target_bhv = max(int(row["bhv_count"]), int(np.ceil(row["occupant_count"] / 50)))
    if target_bhv > row["bhv_count"]:
        alternatives.append(("Increase trained response capacity", "bhv_count", target_bhv))
    if row["alarm_delay_s"] > 60:
        alternatives.append(("Reduce alarm/response delay", "alarm_delay_s", 60))
    if row["communication_channels"] < 3:
        alternatives.append(("Add redundant communication channels", "communication_channels", 3))
    if row["trained_occupants_pct"] < 80:
        alternatives.append(("Increase evacuation familiarity through drills", "trained_occupants_pct", 80.0))
    if int(row["drill_recent"]) == 0:
        alternatives.append(("Run and document a recent evacuation drill", "drill_recent", 1))
    if row["stair_width_m"] < 1.5:
        # Framed only as scenario stress-test; real structural changes require engineering review.
        alternatives.append(("Stress-test wider egress capacity", "stair_width_m", 1.5))

    out: List[Dict[str, object]] = []
    for label, feature, value in alternatives:
        alt = scenario.copy()
        alt.loc[alt.index[0], feature] = value
        alt_pred = float(pipe.predict(alt[FEATURES])[0])
        delta = max(0.0, base_risk - alt_pred)
        out.append({"action": label, "feature": feature, "new_value": value, "risk_reduction": delta})
    return sorted(out, key=lambda x: x["risk_reduction"], reverse=True)
