#!/usr/bin/env python3
"""
OmmAlpha ML Phase 2B - Live Scorer
===================================
Scores leakage-safe LIVE setup rows with the production champions saved by
Phase 2A.

Safeguards:
- exact dataset-version match
- exact per-objective feature-list + SHA256 schema match
- LIVE-only rows and leakage audit
- optional training CSV population guard: only VCP statuses seen by Phase 2A
  are eligible for production inference

New runs apply the validation-selected calibration stored with the model.
Legacy runs remain raw scores; calibration does not guarantee future accuracy.
"""

from __future__ import annotations

import argparse
import hashlib
import json
import os
import sys
from pathlib import Path
from typing import Any

os.environ.setdefault("OMP_NUM_THREADS", "1")
os.environ.setdefault("OPENBLAS_NUM_THREADS", "1")
os.environ.setdefault("MKL_NUM_THREADS", "1")
os.environ.setdefault("NUMEXPR_NUM_THREADS", "1")

import joblib
import numpy as np
import pandas as pd
from probability_calibration import apply_calibration


OBJECTIVES = {
    "trigger": {
        "probability": "trigger_probability",
        "percentile": "trigger_percentile",
        "model": "trigger_model",
    },
    "direct_target": {
        "probability": "target_probability",
        "percentile": "target_percentile",
        "model": "target_model",
    },
    "success": {
        "probability": "conditional_probability",
        "percentile": "conditional_percentile",
        "model": "conditional_model",
    },
}


def parse_args() -> argparse.Namespace:
    p = argparse.ArgumentParser(description="Score OmmAlpha LIVE setups")
    p.add_argument("--csv", required=True, help="LIVE feature CSV exported by MlDatasetService")
    p.add_argument("--run-dir", required=True, help="Completed Phase 2A model run directory")
    p.add_argument("--out", required=True, help="Prediction CSV to create")
    p.add_argument(
        "--training-csv",
        default="",
        help="Phase 2A training CSV used by the production run; used to enforce VCP population compatibility",
    )
    return p.parse_args()


def read_json(path: Path) -> Any:
    with path.open("r", encoding="utf-8") as fh:
        return json.load(fh)


def feature_hash(features: list[str]) -> str:
    return hashlib.sha256("\n".join(features).encode()).hexdigest()


def percentile(values: np.ndarray) -> np.ndarray:
    return (
        pd.Series(values, dtype="float64")
        .rank(method="average", pct=True)
        .mul(100.0)
        .to_numpy()
    )


def validate_live_dataset(df: pd.DataFrame) -> pd.DataFrame:
    required = {
        "dataset_version",
        "signal_id",
        "stock_id",
        "symbol",
        "source",
        "signal_date",
        "vcp_status",
        "leakage_audit_pass",
    }
    missing = sorted(required - set(df.columns))
    if missing:
        raise RuntimeError(f"LIVE CSV missing required columns: {missing}")

    if df.empty:
        raise RuntimeError("LIVE dataset is empty")

    source = df["source"].fillna("").astype(str).str.upper().str.strip()
    non_live = int((source != "LIVE").sum())
    if non_live:
        raise RuntimeError(
            f"Refusing to score mixed/non-LIVE data: {non_live} row(s) are not source=LIVE"
        )

    leakage = pd.to_numeric(df["leakage_audit_pass"], errors="coerce").fillna(0)
    failed = int((leakage != 1).sum())
    if failed:
        raise RuntimeError(
            f"Refusing to score: leakage_audit_pass failed for {failed} LIVE row(s)"
        )

    if df["signal_id"].isna().any():
        raise RuntimeError("Missing signal_id values in LIVE CSV")
    dates = pd.to_datetime(df["signal_date"], errors="coerce")
    if dates.isna().any():
        raise RuntimeError("Invalid or missing signal_date values in LIVE CSV")
    df = df.copy()
    df["signal_date"] = dates.dt.strftime("%Y-%m-%d")
    if df["signal_id"].duplicated().any():
        dupes = int(df["signal_id"].duplicated().sum())
        raise RuntimeError(f"Duplicate signal_id values in LIVE CSV: {dupes}")

    return df.copy()


def training_vcp_statuses(training_csv: Path) -> list[str]:
    if not training_csv.is_file():
        raise RuntimeError(f"Production training CSV is missing: {training_csv}")

    try:
        t = pd.read_csv(
            training_csv,
            usecols=["vcp_status", "leakage_audit_pass"],
            low_memory=False,
        )
    except ValueError as exc:
        raise RuntimeError(
            "Production training CSV does not contain vcp_status/leakage_audit_pass"
        ) from exc

    leakage = pd.to_numeric(t["leakage_audit_pass"], errors="coerce").fillna(0).astype(int)
    statuses = sorted(
        {
            str(v).upper().strip()
            for v in t.loc[leakage == 1, "vcp_status"].dropna().tolist()
            if str(v).strip()
        }
    )
    if not statuses:
        raise RuntimeError("Could not determine VCP statuses from production training CSV")
    return statuses


def load_and_score(
    df: pd.DataFrame,
    run_dir: Path,
    registry: dict[str, Any],
    objective: str,
) -> tuple[np.ndarray, np.ndarray, str, list[str], dict[str, float]]:
    model_info = registry.get("models", {}).get(objective)
    if not isinstance(model_info, dict):
        raise RuntimeError(f"model_registry.json has no '{objective}' model")

    artifact = run_dir / str(model_info.get("artifact", ""))
    features_path = run_dir / str(model_info.get("features", ""))
    metrics_path = run_dir / str(model_info.get("metrics", ""))

    for path in (artifact, features_path, metrics_path):
        if not path.is_file():
            raise RuntimeError(f"{objective}: missing production artifact: {path}")

    features = read_json(features_path)
    if not isinstance(features, list) or not features:
        raise RuntimeError(f"{objective}: invalid/empty features.json")
    features = [str(x) for x in features]

    metrics = read_json(metrics_path)
    expected_hash = str(metrics.get("feature_schema_hash", ""))
    actual_hash = feature_hash(features)
    if not expected_hash or expected_hash != actual_hash:
        raise RuntimeError(f"{objective}: feature schema hash mismatch; scoring stopped")

    missing_features = [c for c in features if c not in df.columns]
    if missing_features:
        raise RuntimeError(
            f"{objective}: LIVE CSV is missing trained features: {missing_features}"
        )

    X = df[features].copy()
    for c in features:
        X[c] = pd.to_numeric(X[c], errors="coerce").replace([np.inf, -np.inf], np.nan)

    if X.isna().all(axis=1).any():
        raise RuntimeError(f"{objective}: LIVE rows have no usable trained features")

    coverage = {c: float(X[c].notna().mean()) for c in features}

    model = joblib.load(artifact)
    if not hasattr(model, "predict_proba"):
        raise RuntimeError(f"{objective}: production champion has no predict_proba()")

    probs = np.asarray(model.predict_proba(X)[:, 1], dtype=float)
    if len(probs) != len(df):
        raise RuntimeError(f"{objective}: prediction row-count mismatch")
    if not np.all(np.isfinite(probs)):
        raise RuntimeError(f"{objective}: non-finite model scores produced")
    if np.any((probs < 0) | (probs > 1)):
        raise RuntimeError(f"{objective}: model scores outside [0, 1]")
    calibration = metrics.get("calibration", {"method": "none"})
    if "calibration" in model_info and model_info["calibration"] != calibration:
        raise RuntimeError(f"{objective}: calibration metadata mismatch")
    probs = apply_calibration(probs, calibration)

    champion = str(model_info.get("champion", metrics.get("champion", "unknown"))).lower()
    ranks = pd.Series(probs, index=df.index).groupby(df["signal_date"]).rank(method="average", pct=True).mul(100).to_numpy()
    return probs, ranks, champion, features, coverage


def main() -> None:
    args = parse_args()
    csv_path = Path(args.csv)
    run_dir = Path(args.run_dir)
    out_path = Path(args.out)
    training_csv = Path(args.training_csv) if args.training_csv else None

    if not csv_path.is_file():
        raise SystemExit(f"Missing LIVE CSV: {csv_path}")
    if not run_dir.is_dir():
        raise SystemExit(f"Missing model run directory: {run_dir}")

    registry_path = run_dir / "model_registry.json"
    if not registry_path.is_file():
        raise SystemExit(f"Missing registry: {registry_path}")

    registry = read_json(registry_path)
    if "deployment" in registry and registry["deployment"].get("eligible") is not True:
        raise RuntimeError("Model run failed validation deployment checks; use an eligible active run")
    if "deployment" in registry and registry["deployment"].get("eligible") is not True:
        raise RuntimeError("Model run failed validation deployment checks; use an eligible active run")
    df = pd.read_csv(csv_path, low_memory=False)
    df = validate_live_dataset(df)

    expected_dataset_version = str(registry.get("dataset_version", ""))
    versions = set(df["dataset_version"].dropna().astype(str).unique())
    if versions != {expected_dataset_version}:
        raise RuntimeError(
            f"Dataset version mismatch. LIVE={sorted(versions)}, model={expected_dataset_version}"
        )

    rows_before_status_filter = int(len(df))
    allowed_statuses: list[str] = []
    if training_csv is not None:
        allowed_statuses = training_vcp_statuses(training_csv)
        live_status = df["vcp_status"].fillna("").astype(str).str.upper().str.strip()
        df = df[live_status.isin(allowed_statuses)].copy()

        if df.empty:
            present = sorted(set(live_status.tolist()) - {""})
            raise RuntimeError(
                "No LIVE rows match the Phase 2A training VCP population. "
                f"Training statuses={allowed_statuses}; LIVE statuses={present}"
            )

    output_cols = [
        "signal_id",
        "stock_id",
        "symbol",
        "source",
        "signal_date",
        "vcp_status",
    ]
    scored = df[output_cols].copy()

    feature_counts: dict[str, int] = {}
    champions: dict[str, str] = {}
    minimum_feature_coverage: dict[str, float] = {}

    for objective, names in OBJECTIVES.items():
        probs, pct, champion, features, coverage = load_and_score(
            df, run_dir, registry, objective
        )
        scored[names["probability"]] = probs
        scored[names["percentile"]] = pct
        scored[names["model"]] = champion
        feature_counts[objective] = len(features)
        champions[objective] = champion
        minimum_feature_coverage[objective] = min(coverage.values()) if coverage else 0.0

    scored.insert(0, "run_key", run_dir.name)
    scored["feature_version"] = expected_dataset_version

    out_path.parent.mkdir(parents=True, exist_ok=True)
    scored.to_csv(out_path, index=False)

    dates = sorted(scored["signal_date"].astype(str).unique().tolist())
    summary = {
        "status": "scored",
        "run_key": run_dir.name,
        "rows": int(len(scored)),
        "rows_before_status_filter": rows_before_status_filter,
        "status_filtered_rows": rows_before_status_filter - int(len(scored)),
        "training_vcp_statuses": allowed_statuses,
        "live_vcp_status_counts": {
            str(k): int(v)
            for k, v in scored["vcp_status"].value_counts().sort_index().to_dict().items()
        },
        "dates": dates,
        "output": str(out_path.resolve()),
        "dataset_version": expected_dataset_version,
        "champions": champions,
        "feature_counts": feature_counts,
        "minimum_live_feature_coverage": minimum_feature_coverage,
        "top_target": (
            scored.sort_values("target_probability", ascending=False)
            .head(10)[
                [
                    "symbol",
                    "vcp_status",
                    "target_probability",
                    "target_percentile",
                    "trigger_probability",
                    "trigger_percentile",
                ]
            ]
            .to_dict(orient="records")
        ),
    }
    print(json.dumps(summary, indent=2, default=str))


if __name__ == "__main__":
    try:
        main()
    except Exception as exc:
        print(json.dumps({"status": "failed", "error": str(exc)}, indent=2), file=sys.stderr)
        raise
