"""
train_hsmm.py
-------------
Offline training script for the duration-aware HSMM regime model used by
HSMM_RegimeEA.mq5 / HSMM_RegimeIndicator.mq5.

Pipeline:
  1. Load an MT5-exported OHLC history CSV.
  2. Build the same three features the MQL5 side computes at runtime
     (realized_vol_dev, efficiency_ratio, return_skew) -- feature parity
     between this script and CHSMMFilter::ComputeFeatures() in
     HSMMFilter.mqh is critical; any drift here silently invalidates the
     whole model, since the live filter would be scoring features on a
     different scale than the one the Gaussians were fit on.
  3. Fit an explicit-duration HSMM via EM (a Baum-Welch analogue extended
     with an explicit per-state duration distribution).
  4. Export every fitted parameter into hsmm_manifest.json, in the exact
     schema CHSMMFilter::LoadManifest() expects.

Usage:
    python train_hsmm.py --csv XAUUSD_M5_history.csv --out hsmm_manifest.json
"""

import argparse
import json
import datetime as dt

import numpy as np
import pandas as pd
from scipy import stats

# ---------------------------------------------------------------------
# Model shape constants -- MUST match the #define block at the top of
# HSMMFilter.mqh (NUM_STATES, NUM_FEATURES, MAX_DURATION,
# FEATURE_LOOKBACK). These are duplicated here rather than shared via
# some cross-language config file because MQL5 has no import mechanism
# for Python constants -- ValidateContract() on the MQL5 side is what
# actually catches a mismatch if the two ever get out of sync.
# ---------------------------------------------------------------------
NUM_STATES = 4
STATE_NAMES = ["TrendUp", "TrendDown", "Range", "HighVolChop"]
NUM_FEATURES = 3
FEATURE_NAMES = ["realized_vol_dev", "efficiency_ratio", "return_skew"]
MAX_DURATION = 60
FEATURE_LOOKBACK = 20
EM_ITERATIONS = 30
EM_TOLERANCE = 1e-4


# ---------------------------------------------------------------------
def load_mt5_csv(path):
    """
    Loads an OHLC history CSV as exported from the MT5 terminal.

    Two MT5-specific quirks are handled here rather than left for pandas'
    defaults to mangle silently:
      - MT5's two export paths disagree on encoding. A CSV written at
        runtime by MQL5's own FILE_CSV (e.g. from an EA/script calling
        FileOpen(..., FILE_CSV)) comes out as UTF-16. A CSV produced via
        the terminal's GUI "Export Bars" button in the Symbols/History
        window, however, comes out as plain UTF-8 (or the system's
        default ANSI codepage) with no BOM at all. A single hard-coded
        encoding would work for one export path and throw a
        UnicodeDecodeError on the other, so several candidate encodings
        are tried in order and the first one that actually parses wins.
      - The column delimiter follows the terminal's Windows regional
        setting (comma OR semicolon depending on locale), so we let
        pandas' python engine sniff it via sep=None rather than hard-
        coding one and having the import silently misparse on a
        different machine.
    """
    last_err = None
    for encoding in ("utf-16", "utf-8-sig", "utf-8", "cp1252"):
        try:
            df = pd.read_csv(path, sep=None, engine="python", encoding=encoding)
            break
        except (UnicodeDecodeError, UnicodeError) as e:
            last_err = e
            df = None
    if df is None:
        raise RuntimeError(
            f"Could not decode {path} with any of utf-16/utf-8-sig/utf-8/cp1252. "
            f"Last error: {last_err}"
        )
    # a UTF-16 BOM sometimes survives as a literal character glued onto
    # the first column name (e.g. "\ufeffTime") -- strip it so column
    # lookups by name don't mysteriously fail, regardless of which
    # encoding above actually matched. The GUI "Export Bars" path also
    # wraps every header in angle brackets (e.g. "<CLOSE>") and splits
    # date and time into two separate columns instead of MQL5's single
    # "time" column -- both get normalized away here so the rest of
    # this script can just ask for "close", "high", "low" without
    # caring which export path produced the file.
    df.columns = [c.replace("\ufeff", "").replace("<", "").replace(">", "") for c in df.columns]
    df = df.rename(columns={c: c.strip().lower() for c in df.columns})

    if "date" in df.columns and "time" in df.columns:
        # GUI export: separate DATE ("2026.01.30") and TIME ("01:05")
        # columns. Concatenated into one string, this sorts correctly
        # lexicographically since both halves are already zero-padded,
        # fixed-width, and in year-month-day / hour-minute order.
        df["time"] = df["date"].astype(str) + " " + df["time"].astype(str)
    elif "date" in df.columns and "time" not in df.columns:
        df = df.rename(columns={"date": "time"})

    sort_col = "time" if "time" in df.columns else df.columns[0]
    df = df.sort_values(sort_col)
    df = df.reset_index(drop=True)
    return df


# ---------------------------------------------------------------------
def compute_features(close, atr):
    """
    Builds the same three features CHSMMFilter::ComputeFeatures()
    computes on the MQL5 side, over a rolling FEATURE_LOOKBACK window.
    Vectorized with numpy for training-time speed across a full history,
    but the underlying formulas are deliberately kept identical,
    term-for-term, to the MQL5 implementation.

    Returns an (N - FEATURE_LOOKBACK, NUM_FEATURES) array; row i
    corresponds to bar index i + FEATURE_LOOKBACK in the original series.
    """
    close = np.asarray(close, dtype=np.float64)
    atr = np.asarray(atr, dtype=np.float64)
    n = len(close)
    log_ret = np.diff(np.log(close))  # log_ret[t] = ln(close[t+1]/close[t])

    n_rows = n - FEATURE_LOOKBACK
    feats = np.zeros((n_rows, NUM_FEATURES), dtype=np.float64)

    for row, t in enumerate(range(FEATURE_LOOKBACK, n)):
        # window of returns ending at bar t: the FEATURE_LOOKBACK returns
        # immediately preceding it, matching CHSMMFilter's
        # close[shift+i]/close[shift+i+1] indexing (newest-first there,
        # but the SET of returns in the window is the same either way).
        window_ret = log_ret[t - FEATURE_LOOKBACK:t]

        realized_sigma = window_ret.std(ddof=0)
        atr_frac = max(atr[t] / close[t], 1e-8)
        feats[row, 0] = (realized_sigma / atr_frac) - 1.0

        net_change = abs(close[t] - close[t - FEATURE_LOOKBACK])
        path_sum = np.abs(np.diff(close[t - FEATURE_LOOKBACK:t + 1])).sum()
        path_sum = max(path_sum, 1e-12)
        feats[row, 1] = net_change / path_sum

        mean_r = window_ret.mean()
        sigma3 = max(realized_sigma ** 3, 1e-12)
        m3 = np.mean((window_ret - mean_r) ** 3)
        feats[row, 2] = m3 / sigma3

    return feats


# ---------------------------------------------------------------------
def gaussian_logpdf_diag(x, mean, var):
    """Diagonal-covariance multivariate Gaussian log-density, summed
    feature-by-feature -- same formula as CHSMMFilter::GaussianLogLik()."""
    var = np.maximum(var, 1e-8)
    return np.sum(-0.5 * np.log(2 * np.pi * var) - 0.5 * (x - mean) ** 2 / var, axis=-1)


# ---------------------------------------------------------------------
def forward_backward_hsmm(feats, trans, dur_pmf, em_mean, em_var):
    """
    Full (non-truncated-for-realtime) explicit-duration forward-backward
    pass over the WHOLE training sequence, used only inside the EM loop.
    This differs from CHSMMFilter's online PredictStep()/UpdateStep() in
    one important way: it also runs a BACKWARD pass, giving smoothed
    (not just filtered) state/duration posteriors -- appropriate for
    offline parameter estimation, where the whole sequence is already
    available, but not something the live EA can do (it only ever sees
    the past).

    Returns:
      gamma:      (T, NUM_STATES)               smoothed P(state=j | all data)
      xi_finish:  (T, NUM_STATES, NUM_STATES)    smoothed P(state j finishes and
                                                  transitions to j' at time t)
      dur_gamma:  (T, NUM_STATES, MAX_DURATION)  smoothed P(state=j, remaining=d)
      total_loglik: float                        true data log-likelihood under the
                                                  CURRENT parameters, log P(O_1:T),
                                                  accumulated from the forward pass's
                                                  own normalizing constants -- this is
                                                  what should actually be tracked for
                                                  EM convergence, not a property of
                                                  gamma (which sums to 1 by
                                                  construction at every t and so
                                                  carries no information about fit
                                                  quality at all).
    """
    T = feats.shape[0]

    # emission log-likelihoods, precomputed once per state for the whole
    # sequence -- this is the single most expensive step per EM
    # iteration, so it is deliberately hoisted out of the T-length loops
    # below rather than recomputed per time step.
    loglik = np.zeros((T, NUM_STATES))
    for j in range(NUM_STATES):
        loglik[:, j] = gaussian_logpdf_diag(feats, em_mean[j], em_var[j])
    lik = np.exp(loglik - loglik.max(axis=1, keepdims=True))

    # forward pass -- same recursion as CHSMMFilter::PredictStep() /
    # UpdateStep(), just kept in an explicit (T, NUM_STATES, MAX_DURATION)
    # array instead of updated in place, because the backward pass below
    # needs every time step's alpha, not just the latest one.
    #
    # This is also a SCALED forward pass in the classic HMM sense: each
    # step's normalizing constant (the total probability mass before
    # dividing it back down to 1) is itself meaningful -- summed in log
    # space across all T steps, it IS the true log-likelihood of the
    # whole observed sequence under the current parameters. Discarding
    # these constants (as an earlier version of this function did) throws
    # away the one number EM convergence should actually be checked
    # against.
    alpha = np.zeros((T, NUM_STATES, MAX_DURATION))
    alpha[0] = 1.0 / (NUM_STATES * MAX_DURATION)
    alpha[0] *= lik[0][:, None]
    c0 = alpha[0].sum()
    alpha[0] /= max(c0, 1e-300)
    # lik[t] = exp(loglik[t] - loglik[t].max()) for numerical stability,
    # which scales every state's likelihood at time t down by a uniform
    # factor of exp(-loglik[t].max()). The true (unscaled) normalizing
    # constant is therefore this step's computed total multiplied back
    # up by exp(+loglik[t].max()) -- i.e. ADD the max back on in log
    # space, not subtract it.
    total_loglik = np.log(max(c0, 1e-300)) + loglik[0].max()

    for t in range(1, T):
        pred = np.zeros((NUM_STATES, MAX_DURATION))
        finishing = alpha[t - 1, :, 0]
        pred[:, :-1] += alpha[t - 1, :, 1:]
        for j in range(NUM_STATES):
            if finishing[j] <= 0:
                continue
            for jp in range(NUM_STATES):
                if jp == j:
                    continue
                outflow = finishing[j] * trans[j, jp]
                if outflow <= 0:
                    continue
                pred[jp, :] += outflow * dur_pmf[jp, :]
        pred *= lik[t][:, None]
        total = pred.sum()
        alpha[t] = pred / total if total > 1e-300 else 1.0 / (NUM_STATES * MAX_DURATION)
        total_loglik += np.log(max(total, 1e-300)) + loglik[t].max()

    # backward pass -- mirror recursion run from T-1 down to 0.
    beta = np.zeros((T, NUM_STATES, MAX_DURATION))
    beta[T - 1] = 1.0
    for t in range(T - 2, -1, -1):
        nxt = beta[t + 1] * lik[t + 1][:, None]
        b = np.zeros((NUM_STATES, MAX_DURATION))
        # continuing (d>0 at t maps to d-1 at t+1)
        b[:, 1:] += nxt[:, :-1]
        # finishing (d==0 at t): probability mass flows out via
        # transition+duration draw, symmetric to the forward pass.
        for j in range(NUM_STATES):
            acc = 0.0
            for jp in range(NUM_STATES):
                if jp == j:
                    continue
                acc += trans[j, jp] * np.sum(dur_pmf[jp, :] * nxt[jp, :])
            b[j, 0] += acc
        s = b.sum()
        beta[t] = b / s if s > 1e-300 else 1.0

    post = alpha * beta
    post /= post.sum(axis=(1, 2), keepdims=True)
    gamma = post.sum(axis=2)                       # (T, NUM_STATES)
    dur_gamma = post                                # (T, NUM_STATES, MAX_DURATION)

    # xi_finish[t, j, jp]: smoothed probability that state j finished
    # (d==0) at time t and the NEXT state is jp -- this is what drives
    # the transition-matrix M-step below.
    xi_finish = np.zeros((T - 1, NUM_STATES, NUM_STATES))
    for t in range(T - 1):
        finishing = alpha[t, :, 0]
        for j in range(NUM_STATES):
            if finishing[j] <= 0:
                continue
            for jp in range(NUM_STATES):
                if jp == j:
                    continue
                mass = finishing[j] * trans[j, jp] * np.sum(dur_pmf[jp, :] * beta[t + 1, jp, :] * lik[t + 1, jp])
                xi_finish[t, j, jp] = mass
        s = xi_finish[t].sum()
        if s > 1e-300:
            xi_finish[t] /= s

    return gamma, xi_finish, dur_gamma, total_loglik


# ---------------------------------------------------------------------
def fit_negative_binomial_moments(weighted_durations, weights):
    """
    Fits a negative-binomial duration distribution to a WEIGHTED sample
    of durations via method-of-moments, then discretizes its pmf over
    1..MAX_DURATION.

    Why parametric smoothing instead of just normalizing the raw
    weighted histogram: with limited training data, the tail bins (long
    durations) get very few effective soft-counts, so a raw histogram's
    tail is noisy and can even contain isolated zero-probability gaps
    that would incorrectly veto a real long-running regime live. Fitting
    a two-parameter negative binomial regularizes the whole curve to a
    smooth, monotonically-plausible shape while still matching the
    empirical mean and variance exactly.
    """
    mean = np.average(weighted_durations, weights=weights)
    var = np.average((weighted_durations - mean) ** 2, weights=weights)
    var = max(var, mean + 1e-6)  # negative binomial requires var > mean
    p = mean / var
    r = mean * p / (1 - p)
    r = max(r, 1e-3)

    k = np.arange(1, MAX_DURATION + 1)
    # shift the standard NB(support 0,1,2,...) so k=1 is the minimum
    # possible sojourn length of one bar.
    pmf = stats.nbinom.pmf(k - 1, r, p)
    pmf = pmf / pmf.sum()
    return pmf


# ---------------------------------------------------------------------
def run_em(feats):
    T = feats.shape[0]

    # --- initialization via k-means (warm start) -------------------
    from scipy.cluster.vq import kmeans2
    centroids, labels = kmeans2(feats, NUM_STATES, minit="++", seed=42)
    # order states by their realized_vol_dev / efficiency_ratio profile
    # so state indices line up with STATE_NAMES in a stable, inspectable
    # way run-to-run rather than an arbitrary k-means label order.
    order = np.lexsort((centroids[:, 1], -centroids[:, 0]))
    remap = {old: new for new, old in enumerate(order)}
    labels = np.array([remap[l] for l in labels])
    centroids = centroids[order]

    em_mean = centroids.copy()
    em_var = np.ones((NUM_STATES, NUM_FEATURES)) * 0.5

    trans = np.full((NUM_STATES, NUM_STATES), 1.0 / (NUM_STATES - 1))
    np.fill_diagonal(trans, 0.0)

    dur_pmf = np.zeros((NUM_STATES, MAX_DURATION))
    for j in range(NUM_STATES):
        dur_pmf[j] = fit_negative_binomial_moments(np.arange(1, MAX_DURATION + 1),
                                                     np.ones(MAX_DURATION))

    prev_loglik = -np.inf
    for it in range(EM_ITERATIONS):
        gamma, xi_finish, dur_gamma, cur_loglik = forward_backward_hsmm(feats, trans, dur_pmf, em_mean, em_var)

        # --- M-step: emissions ---------------------------------
        for j in range(NUM_STATES):
            w = gamma[:, j]
            wsum = w.sum()
            if wsum < 1e-8:
                continue
            em_mean[j] = np.average(feats, axis=0, weights=w)
            diff = feats - em_mean[j]
            em_var[j] = np.average(diff ** 2, axis=0, weights=w)

        # --- M-step: transition matrix (off-diagonal only) -----
        trans_counts = xi_finish.sum(axis=0)
        for j in range(NUM_STATES):
            row_sum = trans_counts[j].sum()
            if row_sum > 1e-8:
                trans[j] = trans_counts[j] / row_sum
            trans[j, j] = 0.0

        # --- M-step: duration distributions ---------------------
        for j in range(NUM_STATES):
            d_values = np.arange(1, MAX_DURATION + 1)
            weights = dur_gamma[:, j, :].sum(axis=0)  # soft count per duration bucket, this state
            if weights.sum() < 1e-6:
                continue
            dur_pmf[j] = fit_negative_binomial_moments(d_values, weights)

        # --- convergence check -----------------------------------
        # cur_loglik is the TRUE data log-likelihood under the
        # parameters used to run the E-step above (i.e. the parameters
        # as of the END of the previous iteration's M-step) -- this is
        # exactly the quantity EM's monotonic-improvement guarantee
        # applies to, and it is expected to increase (or hold flat)
        # every iteration until convergence.
        print(f"EM iteration {it+1:2d}: loglik={cur_loglik:.1f}")
        if abs(cur_loglik - prev_loglik) < EM_TOLERANCE * max(1.0, abs(prev_loglik)):
            print("EM converged.")
            break
        prev_loglik = cur_loglik

    return em_mean, em_var, trans, dur_pmf


# ---------------------------------------------------------------------
def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--csv", required=True, help="MT5-exported OHLC history CSV")
    ap.add_argument("--out", default="hsmm_manifest.json")
    ap.add_argument("--symbol", default="XAUUSD")
    ap.add_argument("--timeframe", default="M5")
    ap.add_argument("--atr-period", type=int, default=14)
    args = ap.parse_args()

    df = load_mt5_csv(args.csv)
    close = df["close"].values

    # a simple Wilder ATR computed here in Python for training purposes
    # only -- it never runs live, since the EA/indicator use MT5's own
    # iATR() handle at runtime, but the FORMULA must match so training
    # features and live features are computed against the same
    # volatility reference.
    high = df["high"].values
    low = df["low"].values
    prev_close = np.roll(close, 1)
    prev_close[0] = close[0]
    tr = np.maximum(high - low, np.maximum(np.abs(high - prev_close), np.abs(low - prev_close)))
    atr = pd.Series(tr).ewm(alpha=1.0 / args.atr_period, adjust=False).mean().values

    feats = compute_features(close, atr)

    feat_mean = feats.mean(axis=0)
    feat_std = feats.std(axis=0)
    feat_std[feat_std < 1e-8] = 1e-8
    feats_z = (feats - feat_mean) / feat_std

    em_mean, em_var, trans, dur_pmf = run_em(feats_z)

    manifest = {
        "schema_version": 1,
        "symbol": args.symbol,
        "timeframe": args.timeframe,
        "num_states": NUM_STATES,
        "state_names": STATE_NAMES,
        "max_duration": MAX_DURATION,
        "feature_names": FEATURE_NAMES,
        "feature_lookback": FEATURE_LOOKBACK,
        "emission_mean": em_mean.tolist(),
        "emission_var": em_var.tolist(),
        "transition_matrix": trans.tolist(),
        "duration_pmf": dur_pmf.tolist(),
        "feature_norm_mean": feat_mean.tolist(),
        "feature_norm_std": feat_std.tolist(),
        "trained_bars": int(feats.shape[0]),
        "training_date": dt.datetime.utcnow().strftime("%Y-%m-%d"),
    }

    with open(args.out, "w", encoding="utf-8") as f:
        json.dump(manifest, f, indent=2)
    print(f"Wrote manifest to {args.out} ({feats.shape[0]} training bars).")


if __name__ == "__main__":
    main()
