#!/usr/bin/env python3
"""Leakage-safe audit and win-rate evidence for Trend Break."""

import hashlib
import json
import math
import os

from neobdm_common import datetime
from trend_break_history import HISTORY_DIR, SCHEMA_VERSION, with_snapshot_hash
from trend_break_outcomes import OUTCOMES_FILE
from tv_enrichment import atomic_write_json


ROOT = os.path.dirname(os.path.abspath(__file__))
BACKTEST_FILE = os.path.join(ROOT, "trend_break_backtest.json")
DATASET_FILE = os.path.join(ROOT, "trend_break_learning_dataset.json")
AUDIT_FILE = os.path.join(ROOT, "trend_break_audit.json")
SCHEMA = "1.0"
TRADING_COST_PCT = 0.30
MIN_SAMPLES = 30
MIN_DAYS = 15

FEATURES = (
    "rsi_1d", "adx_1d", "di_spread_1d", "atr_pct_1d",
    "vwap_distance_pct", "macd_gap_pct", "sma20_distance_pct",
    "sma50_distance_pct", "perf_week_pct", "perf_month_pct",
    "perf_3m_pct", "break_margin_pct", "dist_sma50_pct",
    "close_position_pct", "body_to_range_pct", "upper_wick_pct",
    "volume_ratio_daily", "volume_ratio_intraday", "value_traded",
    "prior_downtrend", "downtrend_proxy", "signal_class",
    "ihsg_change_pct", "ihsg_rsi_1d", "ihsg_adx_1d", "ihsg_regime",
)


def _load_json(path, fallback):
    try:
        with open(path, "r", encoding="utf-8") as handle:
            return json.load(handle)
    except (FileNotFoundError, json.JSONDecodeError, OSError):
        return fallback


def _finite(value, default=0.0):
    try:
        number = float(value)
    except (TypeError, ValueError):
        return default
    return number if math.isfinite(number) else default


def _wilson_low(wins, samples, z=1.96):
    if not samples:
        return None
    p = wins / samples
    denominator = 1 + z * z / samples
    centre = p + z * z / (2 * samples)
    margin = z * math.sqrt(
        p * (1 - p) / samples + z * z / (4 * samples * samples)
    )
    return round(max(0.0, (centre - margin) / denominator) * 100, 1)


def metrics(records):
    labeled = [
        row for row in records
        if row.get("label_status", "LABELED") == "LABELED"
        and row.get("outcome_pct") is not None
    ]
    samples = len(labeled)
    if not samples:
        return {
            "samples": 0,
            "days": 0,
            "wins": 0,
            "losses": 0,
            "winrate_pct": None,
            "net_winrate_pct": None,
            "avg_outcome_pct": None,
            "net_expectancy_pct": None,
            "wilson_low_pct": None,
        }
    outcomes = [_finite(row.get("outcome_pct")) for row in labeled]
    wins = sum(value > 0 for value in outcomes)
    net_wins = sum(value > TRADING_COST_PCT for value in outcomes)
    average = sum(outcomes) / samples
    return {
        "samples": samples,
        "days": len({row.get("date") for row in labeled if row.get("date")}),
        "wins": wins,
        "losses": samples - wins,
        "winrate_pct": round(wins / samples * 100, 1),
        "net_winrate_pct": round(net_wins / samples * 100, 1),
        "avg_outcome_pct": round(average, 4),
        "net_expectancy_pct": round(average - TRADING_COST_PCT, 4),
        "wilson_low_pct": _wilson_low(wins, samples),
    }


def _archive_files(history_dir):
    index = _load_json(os.path.join(history_dir, "index.json"), {})
    return [
        (date, os.path.join(history_dir, f"{date}.json"))
        for date in index.get("available_dates") or []
    ]


def audit_archives(history_dir=HISTORY_DIR):
    issues = []
    payloads = []
    for date, path in _archive_files(history_dir):
        daily = _load_json(path, {})
        payloads.append((date, path, daily))
        snapshots = daily.get("snapshots") or []
        if daily.get("date") != date:
            issues.append({"severity": "ERROR", "date": date, "code": "DATE_MISMATCH"})
        if daily.get("schema_version") != SCHEMA_VERSION:
            issues.append({"severity": "ERROR", "date": date, "code": "SCHEMA_VERSION"})
        if daily.get("snapshot_count") != len(snapshots):
            issues.append({"severity": "ERROR", "date": date, "code": "SNAPSHOT_COUNT"})
        if daily.get("ticker_count") != len(daily.get("tickers") or {}):
            issues.append({"severity": "ERROR", "date": date, "code": "TICKER_COUNT"})
        timestamps = [snapshot.get("timestamp") for snapshot in snapshots]
        if timestamps != sorted(timestamps):
            issues.append({"severity": "ERROR", "date": date, "code": "TIMESTAMP_ORDER"})
        if len(timestamps) != len(set(timestamps)):
            issues.append({"severity": "ERROR", "date": date, "code": "DUPLICATE_TIMESTAMP"})
        expected_archive = hashlib.sha256(
            "".join(
                with_snapshot_hash(snapshot)["content_sha256"]
                for snapshot in snapshots
            ).encode("ascii")
        ).hexdigest()
        if daily.get("archive_sha256") != expected_archive:
            issues.append({"severity": "ERROR", "date": date, "code": "ARCHIVE_HASH"})
        for snapshot in snapshots:
            expected = with_snapshot_hash(snapshot)
            if snapshot.get("content_sha256") != expected["content_sha256"]:
                issues.append({
                    "severity": "ERROR",
                    "date": date,
                    "timestamp": snapshot.get("timestamp"),
                    "code": "SNAPSHOT_HASH",
                })
            signals = snapshot.get("signals") or []
            rejected = snapshot.get("rejected_signals") or []
            if snapshot.get("candidate_count") != len(signals):
                issues.append({"severity": "ERROR", "date": date, "code": "CANDIDATE_COUNT"})
            if int(snapshot.get("rejected_count") or 0) != len(rejected):
                issues.append({"severity": "ERROR", "date": date, "code": "REJECTED_COUNT"})
            tickers = [str(row.get("ticker") or "").upper() for row in signals]
            if len(tickers) != len(set(tickers)):
                issues.append({"severity": "ERROR", "date": date, "code": "DUPLICATE_TICKER"})
            if any(row.get("screening_status") != "KANDIDAT FULL PA" for row in signals):
                issues.append({"severity": "ERROR", "date": date, "code": "INVALID_QUALIFIED_STATUS"})
            if any(row.get("screening_status") != "TIDAK LAYAK DITUNGGU" for row in rejected):
                issues.append({"severity": "ERROR", "date": date, "code": "INVALID_REJECTED_STATUS"})
    errors = sum(issue["severity"] == "ERROR" for issue in issues)
    warnings = sum(issue["severity"] == "WARN" for issue in issues)
    return {
        "status": "PASS" if errors == 0 else "FAIL",
        "files": len(payloads),
        "errors": errors,
        "warnings": warnings,
        "issues": issues,
    }, payloads


def _outcome_map(path):
    mapping = {}
    duplicates = []
    for record in _load_json(path, {"records": []}).get("records") or []:
        key = (
            str(record.get("date") or ""),
            str(record.get("ticker") or "").upper(),
        )
        if key in mapping:
            duplicates.append({"date": key[0], "ticker": key[1]})
        elif all(key):
            mapping[key] = record
    return mapping, duplicates


def build_dataset(history_dir=HISTORY_DIR, outcomes_path=OUTCOMES_FILE):
    integrity, payloads = audit_archives(history_dir)
    outcomes, duplicates = _outcome_map(outcomes_path)
    records = []
    for date, path, daily in payloads:
        first_by_ticker = {}
        for snapshot in daily.get("snapshots") or []:
            for signal in snapshot.get("signals") or []:
                ticker = str(signal.get("ticker") or "").upper()
                context = {
                    "timestamp": snapshot.get("timestamp"),
                    "snapshot_phase": snapshot.get("snapshot_phase"),
                    "market_open": bool(snapshot.get("market_open")),
                    "snapshot_sha256": snapshot.get("content_sha256"),
                    "signal": signal,
                }
                prior = first_by_ticker.get(ticker)
                is_live = context["snapshot_phase"] == "LIVE" and context["market_open"]
                prior_live = bool(
                    prior
                    and prior["snapshot_phase"] == "LIVE"
                    and prior["market_open"]
                )
                if prior is None or (is_live and not prior_live):
                    first_by_ticker[ticker] = context
        for ticker, aggregate in (daily.get("tickers") or {}).items():
            context = first_by_ticker.get(ticker) or {}
            signal = (
                context.get("signal")
                or aggregate.get("first")
                or aggregate.get("latest")
                or {}
            )
            outcome = outcomes.get((date, ticker))
            labeled = bool(outcome and outcome.get("label_d1") in (0, 1))
            features = {
                "price": _finite(signal.get("price")),
                "change_pct": _finite(signal.get("change_pct")),
                "selection_score": _finite(signal.get("selection_score")),
                "signal_class": signal.get("status"),
            }
            features.update(signal.get("technical_features") or {})
            records.append({
                "observation_id": (
                    f"TB:{date}:{ticker}:{context.get('timestamp') or aggregate.get('first_seen_at')}"
                ),
                "date": date,
                "ticker": ticker,
                "feature_timestamp": (
                    context.get("timestamp") or aggregate.get("first_seen_at")
                ),
                "snapshot_phase": context.get("snapshot_phase"),
                "market_open": context.get("market_open"),
                "eligible_for_learning": (
                    context.get("snapshot_phase") == "LIVE"
                    and bool(context.get("market_open"))
                ),
                "lineage": {
                    "archive_file": os.path.relpath(path, ROOT),
                    "snapshot_sha256": context.get("snapshot_sha256"),
                    "outcome_source": (
                        "trend_break_outcomes.json" if labeled else None
                    ),
                },
                "features": features,
                "label_status": "LABELED" if labeled else "PENDING",
                "label": outcome.get("label_d1") if labeled else None,
                "outcome_pct": (
                    outcome.get("d1_close_return_pct") if labeled else None
                ),
                "outcome_definition": (
                    "D+1 close versus first LIVE detection price"
                    if labeled else None
                ),
                "outcomes": {
                    key: outcome.get(key) if outcome else None
                    for key in (
                        "d1_session_date", "d1_close_return_pct",
                        "d1_net_return_pct", "d3_session_date",
                        "d3_close_return_pct", "d3_net_return_pct",
                        "mfe_d3_pct", "mae_d3_pct", "path_2pct",
                        "label_d3",
                    )
                },
            })
    records.sort(key=lambda row: (row["date"], row["feature_timestamp"] or "", row["ticker"]))
    return {
        "schema_version": SCHEMA,
        "generated_at": datetime.now().isoformat(timespec="seconds"),
        "integrity": integrity,
        "outcome_duplicates": duplicates,
        "feature_contract": {
            "unit": "ticker-day pada snapshot LIVE pertama",
            "cost_assumption_pct": TRADING_COST_PCT,
            "guard": "PREMARKET/BREAK/CLOSED tidak eligible; outcome D+1/D+3 bukan fitur",
        },
        "records": records,
    }


def _matches(features, recipe_id):
    signal_class = features.get("signal_class")
    rsi = _finite(features.get("rsi_1d"))
    adx = _finite(features.get("adx_1d"))
    di_spread = _finite(features.get("di_spread_1d"))
    rvol = _finite(features.get("volume_ratio_intraday"))
    break_margin = _finite(features.get("break_margin_pct"), default=-99.0)
    vwap_distance = _finite(features.get("vwap_distance_pct"))
    market_ok = features.get("ihsg_regime") in {"BULLISH", "NEUTRAL"}
    trend_reset = bool(features.get("prior_downtrend"))
    rules = {
        "confirm_pilot": signal_class == "CONFIRM",
        "quality_reset": (
            trend_reset
            and signal_class == "CONFIRM"
            and 45 <= rsi <= 67
            and adx >= 18
            and di_spread > 0
        ),
        "high_wr_lab": (
            trend_reset
            and signal_class == "CONFIRM"
            and -0.5 <= break_margin <= 2.0
            and rvol >= 1.5
            and 45 <= rsi <= 65
            and adx >= 20
            and di_spread > 0
            and vwap_distance >= 0
            and market_ok
        ),
    }
    return bool(rules.get(recipe_id))


def _backtest_records(path=BACKTEST_FILE):
    output = []
    for trade in _load_json(path, {"trades": []}).get("trades") or []:
        return_pct = trade.get("return_pct")
        if return_pct is None:
            continue
        output.append({
            "date": trade.get("signal_date") or trade.get("entry_date"),
            "ticker": trade.get("ticker"),
            "label_status": "LABELED",
            "outcome_pct": _finite(return_pct),
            "features": {
                "signal_class": trade.get("signal_class"),
                "prior_downtrend": True,
                "rsi_1d": 0,
                "adx_1d": 0,
                "di_spread_1d": 0,
                "volume_ratio_intraday": trade.get("vol_ratio"),
                "break_margin_pct": trade.get("break_margin_pct"),
                "vwap_distance_pct": 0,
                "ihsg_regime": "UNKNOWN",
            },
        })
    return output


def _recipes(forward, discovery):
    definitions = (
        ("confirm_pilot", "CONFIRM daily", 55.0, 40.0),
        ("quality_reset", "Confirm + trend reset quality", 60.0, 45.0),
        ("high_wr_lab", "High-WR pre-registered confluence", 85.0, 65.0),
    )
    output = []
    for recipe_id, label, min_wr, min_wilson in definitions:
        forward_selected = [
            row for row in forward
            if row.get("eligible_for_learning")
            and row.get("label_status") == "LABELED"
            and _matches(row.get("features") or {}, recipe_id)
        ]
        # Historical backtest only has enough fields for CONFIRM segmentation.
        discovery_selected = [
            row for row in discovery
            if (
                recipe_id == "confirm_pilot"
                and row.get("features", {}).get("signal_class") == "CONFIRM"
            )
        ]
        forward_metrics = metrics(forward_selected)
        checks = {
            "samples": forward_metrics["samples"] >= MIN_SAMPLES,
            "days": forward_metrics["days"] >= MIN_DAYS,
            "winrate": (
                forward_metrics["winrate_pct"] is not None
                and forward_metrics["winrate_pct"] >= min_wr
            ),
            "net_expectancy": (
                forward_metrics["net_expectancy_pct"] is not None
                and forward_metrics["net_expectancy_pct"] >= 0.30
            ),
            "wilson_low": (
                forward_metrics["wilson_low_pct"] is not None
                and forward_metrics["wilson_low_pct"] >= min_wilson
            ),
        }
        output.append({
            "id": recipe_id,
            "label": label,
            "stage": "LAB" if recipe_id == "high_wr_lab" else "PILOT",
            "discovery_backtest": metrics(discovery_selected),
            "forward_out_of_sample": forward_metrics,
            "promotion_gate": {
                "eligible": all(checks.values()),
                "checks": checks,
                "requirements": {
                    "minimum_samples": MIN_SAMPLES,
                    "minimum_days": MIN_DAYS,
                    "minimum_winrate_pct": min_wr,
                    "minimum_net_expectancy_pct": 0.30,
                    "minimum_wilson_low_pct": min_wilson,
                },
            },
        })
    return output


def _feature_quality(records):
    live = [row for row in records if row.get("eligible_for_learning")]
    coverage = {
        feature: (
            round(sum(
                row.get("features", {}).get(feature) not in (None, "")
                for row in live
            ) / len(live) * 100, 1)
            if live else None
        )
        for feature in FEATURES
    }
    complete = sum(
        all(row.get("features", {}).get(feature) not in (None, "") for feature in FEATURES)
        for row in live
    )
    return {
        "eligible_live_records": len(live),
        "required_feature_count": len(FEATURES),
        "complete_live_records": complete,
        "complete_live_pct": round(complete / len(live) * 100, 1) if live else None,
        "coverage_pct_by_feature": coverage,
        "missing_features": [
            feature for feature, pct in coverage.items()
            if pct is None or pct < 95.0
        ],
    }


def write_audit(
    history_dir=HISTORY_DIR,
    outcomes_path=OUTCOMES_FILE,
    dataset_path=DATASET_FILE,
    audit_path=AUDIT_FILE,
    backtest_path=BACKTEST_FILE,
):
    dataset = build_dataset(history_dir, outcomes_path)
    atomic_write_json(dataset_path, dataset, indent=2)
    records = dataset["records"]
    live = [row for row in records if row.get("eligible_for_learning")]
    labeled = [row for row in live if row.get("label_status") == "LABELED"]
    d3 = [
        row for row in live
        if row.get("outcomes", {}).get("label_d3") in (0, 1)
    ]
    discovery = _backtest_records(backtest_path)
    recipes = _recipes(records, discovery)
    eligible = [
        recipe["id"] for recipe in recipes
        if recipe["promotion_gate"]["eligible"]
    ]
    report = {
        "schema_version": SCHEMA,
        "generated_at": dataset["generated_at"],
        "status": dataset["integrity"]["status"],
        "data_lineage": {
            "snapshot": "trend_break_history/YYYY-MM-DD.json",
            "outcome": "trend_break_outcomes.json (TradingView D+1/D+3)",
            "discovery": "trend_break_backtest.json (Yahoo daily OHLC)",
            "dataset": os.path.basename(dataset_path),
        },
        "integrity": dataset["integrity"],
        "dataset": {
            "forward_records": len(records),
            "forward_live_records": len(live),
            "forward_d1_labeled_records": len(labeled),
            "forward_d3_labeled_records": len(d3),
            "forward_live_pending_records": len(live) - len(labeled),
            "forward_live_dates": len({row["date"] for row in live}),
            "outcome_duplicate_count": len(dataset["outcome_duplicates"]),
            "d1_outcome_coverage_pct": (
                round(len(labeled) / len(live) * 100, 1) if live else None
            ),
        },
        "overall": {
            "discovery_backtest": metrics(discovery),
            "forward_out_of_sample": metrics(labeled),
        },
        "feature_quality": _feature_quality(records),
        "recipes": recipes,
        "learning_readiness": {
            "status": "READY" if eligible else "COLLECTING",
            "ready": bool(eligible),
            "eligible_recipes": eligible,
            "message": (
                "Recipe lolos seluruh forward promotion gate."
                if eligible
                else "Belum ada recipe yang lolos; kumpulkan snapshot LIVE dan outcome D+1."
            ),
        },
        "leakage_guards": [
            "Satu unit learning per ticker per tanggal.",
            "Fitur dibekukan pada snapshot LIVE pertama.",
            "PREMARKET/BREAK/CLOSED tidak masuk promotion gate.",
            "Outcome D+1/D+3 berasal dari sesi setelah tanggal deteksi.",
            "Backtest hanya discovery; kelulusan wajib forward out-of-sample.",
            "High-WR lab menuntut WR >=85%, bukan memaksakan angka 85%.",
        ],
        "warning": (
            "Trend Break adalah screening daily. Win rate tidak mengubah kandidat "
            "menjadi SIAP ENTRY tanpa PA 1D/1H/15m dan trigger candle 5m valid."
        ),
    }
    atomic_write_json(audit_path, report, indent=2)
    return report, dataset


if __name__ == "__main__":
    report, _ = write_audit()
    print(json.dumps({
        "status": report["status"],
        "dataset": report["dataset"],
        "overall": report["overall"],
        "learning_readiness": report["learning_readiness"],
    }, indent=2))
