"""verify.py -- DexIQ Signal Track Record skeptic-grade reproducibility tool.

Reads a manifest.json (as published at https://track.dexiq.io/manifest.json),
recomputes the 5 top-line metrics, and prints them. NO DexIQ imports -- the
script depends only on Python stdlib + numpy + scipy.

Usage:
    python verify.py /path/to/manifest.json

Exit codes:
    0 -- success, metrics printed
    1 -- I/O or parse error
"""
import json
import math
import sys
from typing import Optional

import numpy as np
from scipy.stats import norm  # type: ignore


EULER_GAMMA = 0.5772156649015329


def compute_sharpe(returns, periods_per_year: int = 365 * 6) -> float:
    """Annualised Sharpe from a return series. Periods/year default 2190 (4h)."""
    n = len(returns)
    if n < 2:
        return 0.0
    mean_r = sum(returns) / n
    variance = sum((r - mean_r) ** 2 for r in returns) / (n - 1)
    std_r = math.sqrt(variance) if variance > 0 else 0.0
    if std_r == 0.0:
        return 0.0
    return (mean_r / std_r) * math.sqrt(periods_per_year)


def deflated_sharpe(observed_sharpe: float, n_trials: int,
                    variance_of_sharpe_estimates: float, T: int) -> float:
    """Bailey-Lopez de Prado deflated Sharpe p-value."""
    if n_trials < 1 or T < 1:
        raise ValueError("n_trials and T must be >= 1")
    if n_trials == 1:
        expected_max_sr = 0.0
    else:
        e = math.e
        z1 = norm.ppf(1.0 - 1.0 / n_trials)
        z2 = norm.ppf(1.0 - 1.0 / (n_trials * e))
        expected_max_sr = math.sqrt(variance_of_sharpe_estimates) * (
            (1 - EULER_GAMMA) * z1 + EULER_GAMMA * z2
        )
    std_sr = math.sqrt(variance_of_sharpe_estimates / T)
    if std_sr == 0.0:
        return 1.0 if observed_sharpe > expected_max_sr else 0.0
    return float(norm.cdf((observed_sharpe - expected_max_sr) / std_sr))


def main(argv):
    if len(argv) < 2:
        print("usage: python verify.py <manifest.json>", file=sys.stderr)
        return 1
    try:
        with open(argv[1], "r", encoding="utf-8") as fh:
            rows = json.load(fh)
    except (OSError, json.JSONDecodeError) as exc:
        print(f"could not read manifest: {exc}", file=sys.stderr)
        return 1

    # Settled rows have realized_return != null.
    settled = [r for r in rows if r.get("realized_return") is not None]

    # Per-asset win rate excludes flat_skip.
    per_asset: dict = {}
    for r in settled:
        if r.get("exit_type") == "flat_skip":
            continue
        bucket = per_asset.setdefault(r["asset"], [0, 0])
        bucket[1] += 1
        if r["realized_return"] > 0:
            bucket[0] += 1
    win_rate = {a: (w / t) for a, (w, t) in per_asset.items() if t > 0}

    # Cumulative compounded
    returns = [float(r["realized_return"]) for r in
               sorted(settled, key=lambda x: x["signal_time"])]
    equity = 1.0
    for r in returns:
        equity *= (1.0 + r)
    cumulative = equity - 1.0 if returns else 0.0

    # Max drawdown
    if returns:
        eq = np.cumprod(1.0 + np.array(returns, dtype=float))
        peaks = np.maximum.accumulate(eq)
        dd = eq / peaks - 1.0
        mdd = float(np.min(dd))
    else:
        mdd = 0.0

    # Deflated Sharpe
    engines = sorted({r["engine_source"] for r in rows})
    n_trials = max(len(engines), 1)
    if len(returns) >= 2:
        observed = compute_sharpe(returns)
        try:
            dsr = deflated_sharpe(observed, n_trials, 0.5, len(returns))
        except ValueError:
            dsr = None
    else:
        dsr = None

    print("DexIQ Signal Track Record -- verify.py")
    print("=" * 50)
    print(f"  Total signals       : {len(rows)}")
    print(f"  Settled signals     : {len(settled)}")
    print(f"  Win rate by asset   : {win_rate}")
    print(f"  Cumulative return   : {cumulative * 100:+.4f}%")
    print(f"  Max drawdown        : {mdd * 100:+.4f}%")
    if dsr is None:
        print("  Deflated Sharpe p   : (insufficient data)")
    else:
        print(f"  Deflated Sharpe p   : {dsr:.4f}")
    print("=" * 50)
    return 0


if __name__ == "__main__":
    sys.exit(main(sys.argv))
