#!/usr/bin/env python3
"""
standalone_bayes_gp_cli.py

Standalone Bayesian optimization / adaptive learning CLI.

- Does NOT import tkbo or tklib.
- Reads CSV / Excel.
- Infers target and feature columns using tkbo-like column rules.
- Splits observed rows into train/validation.
- Fits a Gaussian Process regression model.
- Saves and loads trained models with joblib.
- Scores unevaluated candidates by EI / PI / UCB / LCB / entropy / stein-lite.

Examples
--------
Read input summary:
    python standalone_bayes_gp_cli.py --mode read --infile data.xlsx

Train, validate, predict all rows, and suggest next candidates:
    python standalone_bayes_gp_cli.py --mode ask --infile data.xlsx --n-points 3 --outfile result.xlsx --model-file gp_model.joblib

Reuse a saved model without refitting:
    python standalone_bayes_gp_cli.py --mode predict --infile data.xlsx --load-model gp_model.joblib --outfile pred.xlsx

Explicit target/features:
    python standalone_bayes_gp_cli.py --mode ask --infile data.csv --target "max:y" --features x1,x2,x3
"""

from __future__ import annotations

import argparse
import json
import math
import re
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Iterable, Optional, Sequence

import joblib
import numpy as np
import pandas as pd
from sklearn.gaussian_process import GaussianProcessRegressor
from sklearn.gaussian_process.kernels import ConstantKernel, RBF, WhiteKernel
from sklearn.metrics import 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 StandardScaler


# =============================================================================
# tkbo-like base utilities
# =============================================================================


@dataclass
class BOResult:
    indices: np.ndarray
    scores: Optional[np.ndarray] = None
    X: Optional[np.ndarray] = None


@dataclass
class DataBundle:
    df: pd.DataFrame
    target: str
    target_label: str
    target_mode: str
    target_value: Optional[float]
    features: list[str]
    X: np.ndarray
    y_original: np.ndarray
    y_bo: np.ndarray
    observed_mask: np.ndarray
    valid_feature_mask: np.ndarray


def as_2d_array(X: Any, name: str = "X") -> np.ndarray:
    arr = np.asarray(X, dtype=float)
    if arr.ndim == 1:
        arr = arr.reshape(-1, 1)
    if arr.ndim != 2:
        raise ValueError(f"{name} must be 1D or 2D array-like, got shape={arr.shape}")
    return arr


def as_1d_array(y: Any, name: str = "y") -> np.ndarray:
    arr = np.asarray(y, dtype=float).reshape(-1)
    if arr.ndim != 1:
        raise ValueError(f"{name} must be 1D array-like")
    return arr


def read_table(path: str | Path, sheet_name: str | int = 0) -> pd.DataFrame:
    path = Path(path)
    if not path.exists():
        raise FileNotFoundError(f"Input file not found: {path}")
    if path.suffix.lower() in {".xlsx", ".xlsm", ".xls"}:
        return pd.read_excel(path, sheet_name=sheet_name)
    return pd.read_csv(path)


def write_table(df: pd.DataFrame, path: str | Path) -> None:
    path = Path(path)
    path.parent.mkdir(parents=True, exist_ok=True)
    if path.suffix.lower() in {".xlsx", ".xlsm", ".xls"}:
        df.to_excel(path, index=False)
    else:
        df.to_csv(path, index=False)


def parse_target_column(name: str) -> tuple[str, str, Optional[float]]:
    """Return (label_without_prefix, mode, target_value).

    Supported target prefixes:
      max:xxx, max1:xxx, t:xxx, o:xxx -> maximize
      min:xxx, min1:xxx                  -> minimize, converted to -y
      =1.23:xxx                          -> target value, converted to -(y-v)^2
    """
    s = str(name)
    m = re.match(r"^=([+-]?\d*\.?\d+(?:[eE][+-]?\d+)?):(.*)$", s)
    if m:
        return m.group(2), "value", float(m.group(1))

    m = re.match(r"^(max\d*|t\d*|o\d*):(.*)$", s, flags=re.IGNORECASE)
    if m:
        return m.group(2), "max", None

    m = re.match(r"^(min\d*):(.*)$", s, flags=re.IGNORECASE)
    if m:
        return m.group(2), "min", None

    return s, "max", None


def infer_columns(
    df: pd.DataFrame,
    target: str | None = None,
    features: str | list[str] | None = None,
) -> tuple[str, list[str]]:
    columns = list(df.columns)
    if target is None:
        candidates = [
            c
            for c in columns
            if re.match(
                r"^(target|t\d*:|o\d*:|max\d*:|min\d*:|=[+-]?\d*\.?\d+(?:[eE][+-]?\d+)?:)",
                str(c),
                flags=re.IGNORECASE,
            )
        ]
        target = candidates[0] if candidates else columns[0]

    if isinstance(features, str):
        features = [s.strip() for s in features.split(",") if s.strip()]

    if features is None:
        features = [
            c
            for c in columns
            if c != target and not str(c).startswith("-") and pd.api.types.is_numeric_dtype(df[c])
        ]

    missing = [c for c in [target, *features] if c not in df.columns]
    if missing:
        raise ValueError(f"Columns not found in input: {missing}")
    if len(features) == 0:
        raise ValueError("No numeric feature columns were found. Use --features explicitly.")
    return str(target), [str(c) for c in features]


def transform_target_for_bo(y: np.ndarray, mode: str, target_value: Optional[float]) -> np.ndarray:
    y = np.asarray(y, dtype=float).copy()
    if mode == "min":
        return -y
    if mode == "value":
        if target_value is None:
            raise ValueError("target_value is required for value mode")
        return -((y - float(target_value)) ** 2)
    return y


def inverse_mean_for_display(y_bo_mean: np.ndarray, mode: str, target_value: Optional[float]) -> np.ndarray:
    """Best-effort inverse for display only.

    For '=value:' mode, the exact inverse is not unique. We keep BO-scale mean.
    """
    y_bo_mean = np.asarray(y_bo_mean, dtype=float)
    if mode == "min":
        return -y_bo_mean
    return y_bo_mean


def split_observed_candidates(df: pd.DataFrame, target: str, features: list[str]) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
    X_df = df[features].apply(pd.to_numeric, errors="coerce")
    y = pd.to_numeric(df[target], errors="coerce").to_numpy(dtype=float)
    X = X_df.to_numpy(dtype=float)
    observed_mask = ~np.isnan(y)
    valid_feature_mask = ~np.any(np.isnan(X), axis=1)
    return X, y, observed_mask, valid_feature_mask


def load_data(path: str | Path, target: str | None = None, features: str | list[str] | None = None, sheet_name: str | int = 0) -> DataBundle:
    df = read_table(path, sheet_name=sheet_name)
    target, features = infer_columns(df, target=target, features=features)
    label, mode, target_value = parse_target_column(target)
    X, y_original, observed_mask, valid_feature_mask = split_observed_candidates(df, target, features)
    y_bo = transform_target_for_bo(y_original, mode, target_value)
    return DataBundle(
        df=df,
        target=target,
        target_label=label,
        target_mode=mode,
        target_value=target_value,
        features=features,
        X=X,
        y_original=y_original,
        y_bo=y_bo,
        observed_mask=observed_mask,
        valid_feature_mask=valid_feature_mask,
    )


# =============================================================================
# Acquisition functions
# =============================================================================


def _normal_pdf(z: np.ndarray) -> np.ndarray:
    return np.exp(-0.5 * z * z) / math.sqrt(2.0 * math.pi)


def _normal_cdf(z: np.ndarray) -> np.ndarray:
    # numpy does not always expose erf. math.erf vectorization is enough here.
    return 0.5 * (1.0 + np.vectorize(math.erf)(z / math.sqrt(2.0)))


def acquisition_score(
    name: str,
    X: np.ndarray,
    mean: np.ndarray,
    std: np.ndarray,
    y_observed: np.ndarray,
    observed_mask: Optional[np.ndarray] = None,
    model: Any = None,
    xi: float = 0.0,
    kappa: float = 2.0,
    eps: float = 1.0e-5,
    std_weight: float = 1.0,
    grad_weight: float = 1.0,
) -> np.ndarray:
    name = name.lower()
    mean = np.asarray(mean, dtype=float).reshape(-1)
    std = np.maximum(np.asarray(std, dtype=float).reshape(-1), 1.0e-300)
    y_observed = np.asarray(y_observed, dtype=float).reshape(-1)
    y_best = np.nanmax(y_observed) if len(y_observed) else 0.0

    if name in {"ei", "expected_improvement"}:
        improvement = mean - y_best - xi
        z = improvement / std
        score = improvement * _normal_cdf(z) + std * _normal_pdf(z)
    elif name in {"pi", "probability_improvement"}:
        z = (mean - y_best - xi) / std
        score = _normal_cdf(z)
    elif name == "ucb":
        score = mean + kappa * std
    elif name == "lcb":
        score = -mean + kappa * std
    elif name == "entropy":
        score = np.log(std)
    elif name in {"stein", "stein-lite", "stein_lite"}:
        if model is None:
            score = std.copy()
        else:
            grad = np.zeros_like(X, dtype=float)
            for j in range(X.shape[1]):
                Xp = X.copy(); Xp[:, j] += eps
                Xm = X.copy(); Xm[:, j] -= eps
                mp = np.asarray(model.predict(Xp), dtype=float).reshape(-1)
                mm = np.asarray(model.predict(Xm), dtype=float).reshape(-1)
                grad[:, j] = (mp - mm) / (2.0 * eps)
            score = std_weight * std + grad_weight * std * np.linalg.norm(grad, axis=1)
    elif name == "random":
        rng = np.random.default_rng()
        score = rng.random(len(mean))
    else:
        raise ValueError(f"Unknown acquisition: {name}")

    if observed_mask is not None:
        score = score.copy()
        score[np.asarray(observed_mask, dtype=bool)] = -np.inf
    return score


def top_k_from_score(score: np.ndarray, k: int) -> np.ndarray:
    score = np.asarray(score, dtype=float)
    valid = np.where(np.isfinite(score))[0]
    if len(valid) == 0:
        return np.array([], dtype=int)
    order = valid[np.argsort(score[valid])[::-1]]
    return order[: max(0, min(int(k), len(order)))]


# =============================================================================
# Model construction, training, evaluation
# =============================================================================


def build_gpr_model(
    standardize: int = 1,
    alpha: float = 1.0e-10,
    normalize_y: int = 1,
    n_restarts_optimizer: int = 10,
    random_seed: int | None = None,
    length_scale: float = 1.0,
    noise_level: float | None = None,
) -> Pipeline | GaussianProcessRegressor:
    kernel = ConstantKernel(1.0, (1.0e-3, 1.0e3)) * RBF(
        length_scale=length_scale,
        length_scale_bounds=(1.0e-4, 1.0e4),
    )
    if noise_level is not None:
        kernel = kernel + WhiteKernel(noise_level=noise_level, noise_level_bounds=(1.0e-12, 1.0e1))

    gpr = GaussianProcessRegressor(
        kernel=kernel,
        alpha=alpha,
        normalize_y=bool(normalize_y),
        n_restarts_optimizer=n_restarts_optimizer,
        random_state=random_seed,
    )
    if standardize:
        return Pipeline([("scaler", StandardScaler()), ("gpr", gpr)])
    return gpr


def predict_mean_std(model: Any, X: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
    X = as_2d_array(X)
    mean, std = model.predict(X, return_std=True)
    return np.asarray(mean, dtype=float).reshape(-1), np.asarray(std, dtype=float).reshape(-1)


def fit_or_load_model(args: argparse.Namespace, X_train: np.ndarray, y_train: np.ndarray) -> Any:
    if args.load_model:
        model = joblib.load(args.load_model)
        return model

    model = build_gpr_model(
        standardize=args.standardize,
        alpha=args.alpha,
        normalize_y=args.normalize_y,
        n_restarts_optimizer=args.n_restarts_optimizer,
        random_seed=args.random_seed,
        length_scale=args.length_scale,
        noise_level=args.noise_level,
    )
    model.fit(X_train, y_train)
    return model


def evaluate_model(model: Any, X_val: np.ndarray, y_val: np.ndarray) -> dict[str, float]:
    if len(y_val) == 0:
        return {}
    y_pred = np.asarray(model.predict(X_val), dtype=float).reshape(-1)
    rmse = float(np.sqrt(mean_squared_error(y_val, y_pred)))
    mae = float(mean_absolute_error(y_val, y_pred))
    metrics = {"rmse_bo": rmse, "mae_bo": mae}
    if len(y_val) >= 2:
        metrics["r2_bo"] = float(r2_score(y_val, y_pred))
    return metrics


def default_outfile(infile: str | Path, suffix: str = "-predict.xlsx") -> str:
    p = Path(infile)
    return str(p.with_name(p.stem + suffix))


def default_model_file(infile: str | Path) -> str:
    p = Path(infile)
    return str(p.with_name(p.stem + "-model.joblib"))


def build_output_frame(bundle: DataBundle, mean: np.ndarray, std: np.ndarray, score: Optional[np.ndarray] = None) -> pd.DataFrame:
    out = bundle.df.copy()
    out["__observed__"] = bundle.observed_mask
    out["__valid_features__"] = bundle.valid_feature_mask
    out["__pred_mean_bo__"] = mean
    out["__pred_std_bo__"] = std
    out["__pred_mean_display__"] = inverse_mean_for_display(mean, bundle.target_mode, bundle.target_value)
    if score is not None:
        out["__acquisition_score__"] = score
        out["__suggest_rank__"] = np.nan
        finite = np.where(np.isfinite(score))[0]
        order = finite[np.argsort(score[finite])[::-1]]
        for rank, idx in enumerate(order, 1):
            out.loc[out.index[idx], "__suggest_rank__"] = rank
    return out


def train_validation_indices(
    candidate_indices: np.ndarray,
    y: np.ndarray,
    validation_fraction: float,
    random_seed: int | None,
) -> tuple[np.ndarray, np.ndarray]:
    candidate_indices = np.asarray(candidate_indices, dtype=int)
    if validation_fraction <= 0 or len(candidate_indices) < 3:
        return candidate_indices, np.array([], dtype=int)
    n_val = max(1, int(round(len(candidate_indices) * validation_fraction)))
    n_val = min(n_val, len(candidate_indices) - 1)
    train_idx, val_idx = train_test_split(
        candidate_indices,
        test_size=n_val,
        random_state=random_seed,
        shuffle=True,
    )
    return np.asarray(train_idx, dtype=int), np.asarray(val_idx, dtype=int)


# =============================================================================
# CLI modes
# =============================================================================


def mode_read(args: argparse.Namespace) -> None:
    bundle = load_data(args.infile, target=args.target, features=args.features, sheet_name=args.sheet_name)
    observed = bundle.observed_mask & bundle.valid_feature_mask
    candidates = (~bundle.observed_mask) & bundle.valid_feature_mask
    invalid = ~bundle.valid_feature_mask
    print(f"Input: {args.infile}")
    print(f"Rows: {len(bundle.df)}")
    print(f"Target: {bundle.target}  mode={bundle.target_mode}  value={bundle.target_value}")
    print(f"Features: {bundle.features}")
    print(f"Observed rows with valid features: {int(observed.sum())}")
    print(f"Candidate rows with blank target and valid features: {int(candidates.sum())}")
    print(f"Rows skipped due to NaN in features: {int(invalid.sum())}")
    print(bundle.df.head(args.head).to_string(index=False))


def mode_train_predict_ask(args: argparse.Namespace) -> None:
    bundle = load_data(args.infile, target=args.target, features=args.features, sheet_name=args.sheet_name)

    observed_indices = np.where(bundle.observed_mask & bundle.valid_feature_mask)[0]
    if len(observed_indices) == 0 and not args.load_model:
        raise ValueError("No observed rows are available for training. Fill at least one target value or use --load-model.")

    train_idx, val_idx = train_validation_indices(
        observed_indices,
        bundle.y_bo,
        validation_fraction=args.validation_fraction,
        random_seed=args.random_seed,
    )
    X_train = bundle.X[train_idx]
    y_train = bundle.y_bo[train_idx]
    X_val = bundle.X[val_idx]
    y_val = bundle.y_bo[val_idx]

    model = fit_or_load_model(args, X_train, y_train)

    metrics = evaluate_model(model, X_val, y_val)
    if metrics:
        print("Validation metrics on BO-scale target:")
        for k, v in metrics.items():
            print(f"  {k}: {v:.6g}")
    else:
        print("Validation split is empty. Use --validation-fraction > 0 with at least 3 observed rows to evaluate validation metrics.")

    valid_predict_mask = bundle.valid_feature_mask
    mean = np.full(len(bundle.df), np.nan, dtype=float)
    std = np.full(len(bundle.df), np.nan, dtype=float)
    mean_valid, std_valid = predict_mean_std(model, bundle.X[valid_predict_mask])
    mean[valid_predict_mask] = mean_valid
    std[valid_predict_mask] = std_valid

    score = None
    selected = np.array([], dtype=int)
    if args.mode in {"ask", "train"}:
        observed_for_ask = bundle.observed_mask | (~bundle.valid_feature_mask)
        y_observed = bundle.y_bo[observed_indices]
        score = acquisition_score(
            args.acquisition,
            X=bundle.X,
            mean=mean,
            std=std,
            y_observed=y_observed,
            observed_mask=observed_for_ask,
            model=model,
            xi=args.xi,
            kappa=args.kappa,
            eps=args.stein_eps,
            std_weight=args.stein_std_weight,
            grad_weight=args.stein_grad_weight,
        )
        selected = top_k_from_score(score, args.n_points)

        print("\nSuggested candidates:")
        for rank, idx in enumerate(selected, 1):
            excel_row = idx + 2
            x_dict = {col: bundle.df.iloc[idx][col] for col in bundle.features}
            print(
                f"  #{rank}: index={idx}, Excel row={excel_row}, "
                f"mean_bo={mean[idx]:.6g}, std={std[idx]:.6g}, score={score[idx]:.6g}, X={x_dict}"
            )

    if args.save_model:
        model_file = args.model_file or default_model_file(args.infile)
        payload = {
            "model": model,
            "target": bundle.target,
            "target_label": bundle.target_label,
            "target_mode": bundle.target_mode,
            "target_value": bundle.target_value,
            "features": bundle.features,
            "metadata": {
                "acquisition": args.acquisition,
                "standardize": args.standardize,
                "train_indices": train_idx.tolist(),
                "validation_indices": val_idx.tolist(),
                "metrics": metrics,
            },
        }
        joblib.dump(payload, model_file)
        print(f"\nSaved model: {model_file}")

    outfile = args.outfile or default_outfile(args.infile)
    if args.save:
        out = build_output_frame(bundle, mean=mean, std=std, score=score)
        if selected.size:
            out["__selected__"] = False
            for idx in selected:
                out.loc[out.index[idx], "__selected__"] = True
        write_table(out, outfile)
        print(f"Saved output: {outfile}")


def mode_predict_with_saved_model(args: argparse.Namespace) -> None:
    if not args.load_model:
        raise ValueError("--mode predict requires --load-model")
    payload = joblib.load(args.load_model)
    model = payload["model"] if isinstance(payload, dict) and "model" in payload else payload

    bundle = load_data(
        args.infile,
        target=args.target or (payload.get("target") if isinstance(payload, dict) else None),
        features=args.features or (payload.get("features") if isinstance(payload, dict) else None),
        sheet_name=args.sheet_name,
    )
    valid_predict_mask = bundle.valid_feature_mask
    mean = np.full(len(bundle.df), np.nan, dtype=float)
    std = np.full(len(bundle.df), np.nan, dtype=float)
    mean_valid, std_valid = predict_mean_std(model, bundle.X[valid_predict_mask])
    mean[valid_predict_mask] = mean_valid
    std[valid_predict_mask] = std_valid

    outfile = args.outfile or default_outfile(args.infile)
    out = build_output_frame(bundle, mean=mean, std=std, score=None)
    write_table(out, outfile)
    print(f"Loaded model: {args.load_model}")
    print(f"Saved prediction: {outfile}")


def mode_tell(args: argparse.Namespace) -> None:
    """Update target values in the input table.

    --tell-values format:
        index:value,index:value

    Index is zero-based DataFrame index. Excel row is index + 2.
    """
    if not args.tell_values:
        raise ValueError("--mode tell requires --tell-values, e.g. --tell-values 10:3.14,11:2.71")
    bundle = load_data(args.infile, target=args.target, features=args.features, sheet_name=args.sheet_name)
    df = bundle.df.copy()
    for pair in args.tell_values.split(","):
        idx_s, val_s = pair.split(":", 1)
        idx = int(idx_s.strip())
        val = float(val_s.strip())
        if idx < 0 or idx >= len(df):
            raise IndexError(f"tell index out of range: {idx}")
        df.loc[df.index[idx], bundle.target] = val
    outfile = args.outfile or default_outfile(args.infile, suffix="-tell.xlsx")
    write_table(df, outfile)
    print(f"Saved updated table: {outfile}")


# =============================================================================
# argparse
# =============================================================================


def build_parser() -> argparse.ArgumentParser:
    p = argparse.ArgumentParser(description="Standalone Gaussian-process Bayesian optimization CLI")
    p.add_argument("--mode", default="ask", choices=["read", "train", "ask", "predict", "tell"])
    p.add_argument("--infile", required=True)
    p.add_argument("--outfile", default="")
    p.add_argument("--sheet-name", default=0)
    p.add_argument("--target", default=None)
    p.add_argument("--features", default=None, help="Comma-separated feature columns")

    p.add_argument("--acquisition", default="ei", choices=["ei", "pi", "ucb", "lcb", "entropy", "stein", "random"])
    p.add_argument("--n-points", type=int, default=1)
    p.add_argument("--validation-fraction", type=float, default=0.2)
    p.add_argument("--standardize", type=int, default=1, choices=[0, 1])
    p.add_argument("--normalize-y", type=int, default=1, choices=[0, 1])
    p.add_argument("--random-seed", type=int, default=None)

    p.add_argument("--alpha", type=float, default=1.0e-10)
    p.add_argument("--length-scale", type=float, default=1.0)
    p.add_argument("--noise-level", type=float, default=None)
    p.add_argument("--n-restarts-optimizer", type=int, default=10)

    p.add_argument("--xi", type=float, default=0.0, help="EI/PI exploration offset")
    p.add_argument("--kappa", type=float, default=2.0, help="UCB/LCB exploration weight")
    p.add_argument("--stein-eps", type=float, default=1.0e-5)
    p.add_argument("--stein-std-weight", type=float, default=1.0)
    p.add_argument("--stein-grad-weight", type=float, default=1.0)

    p.add_argument("--save", type=int, default=1, choices=[0, 1])
    p.add_argument("--save-model", type=int, default=1, choices=[0, 1])
    p.add_argument("--model-file", default="")
    p.add_argument("--load-model", default="")
    p.add_argument("--tell-values", default="")
    p.add_argument("--head", type=int, default=5)
    return p


def main(argv: Optional[Sequence[str]] = None) -> int:
    parser = build_parser()
    args = parser.parse_args(argv)

    # argparse parses --sheet-name as str. Convert integer-like sheet names.
    if isinstance(args.sheet_name, str) and args.sheet_name.isdigit():
        args.sheet_name = int(args.sheet_name)

    try:
        if args.mode == "read":
            mode_read(args)
        elif args.mode == "predict":
            mode_predict_with_saved_model(args)
        elif args.mode == "tell":
            mode_tell(args)
        else:
            mode_train_predict_ask(args)
    except Exception as exc:
        print(f"ERROR: {exc}")
        if hasattr(args, "debug") and args.debug:
            raise
        return 1
    return 0


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