#!/usr/bin/env python3
"""Dependency-free reference calculations for the Sharpe-ratio chapter."""

from __future__ import annotations

import math
import statistics
from collections.abc import Sequence


NORMAL = statistics.NormalDist()
EULER_MASCHERONI = 0.5772156649015329


def iid_annualized_sharpe_se(
    annualized_sharpe: float, observations: int, periods_per_year: int
) -> float:
    """IID-normal approximation, expressed in annualized Sharpe units."""
    if observations < 2 or periods_per_year < 1:
        raise ValueError("observations >= 2 and periods_per_year >= 1 are required")
    return math.sqrt(
        (periods_per_year + annualized_sharpe**2 / 2.0) / observations
    )


def probabilistic_sharpe_ratio(
    annualized_sharpe: float,
    benchmark_annualized_sharpe: float,
    observations: int,
    skewness: float,
    raw_kurtosis: float,
    periods_per_year: int,
) -> float:
    """Bailey-Lopez de Prado PSR; kurtosis is raw (Gaussian = 3)."""
    if observations < 2 or periods_per_year < 1:
        raise ValueError("observations >= 2 and periods_per_year >= 1 are required")
    scale = math.sqrt(periods_per_year)
    sr = annualized_sharpe / scale
    benchmark = benchmark_annualized_sharpe / scale
    variance_term = 1.0 - skewness * sr + (raw_kurtosis - 1.0) * sr**2 / 4.0
    if variance_term <= 0:
        raise ValueError("PSR variance term must be positive")
    z = (sr - benchmark) * math.sqrt(observations - 1) / math.sqrt(variance_term)
    return NORMAL.cdf(z)


def expected_maximum_sharpe_under_null(
    trial_sharpes: Sequence[float],
) -> float:
    """DSR's zero-mean-null benchmark from Bailey-Lopez de Prado (2014)."""
    trials = len(trial_sharpes)
    if trials < 1:
        raise ValueError("at least one trial Sharpe estimate is required")
    if trials == 1:
        return 0.0
    sigma = statistics.stdev(trial_sharpes)
    upper = NORMAL.inv_cdf(1.0 - 1.0 / trials)
    upper_e = NORMAL.inv_cdf(1.0 - 1.0 / (trials * math.e))
    return sigma * (
        (1.0 - EULER_MASCHERONI) * upper + EULER_MASCHERONI * upper_e
    )


def deflated_sharpe_ratio(
    selected_annualized_sharpe: float,
    trial_annualized_sharpes: Sequence[float],
    observations: int,
    skewness: float,
    raw_kurtosis: float,
    periods_per_year: int,
) -> float:
    """DSR is PSR evaluated against the trials' expected maximum Sharpe."""
    if not trial_annualized_sharpes:
        raise ValueError("at least one trial Sharpe estimate is required")
    if not math.isclose(
        selected_annualized_sharpe,
        max(trial_annualized_sharpes),
        rel_tol=1e-12,
        abs_tol=1e-12,
    ):
        raise ValueError("selected Sharpe must be the maximum of the supplied trials")
    benchmark = expected_maximum_sharpe_under_null(trial_annualized_sharpes)
    return probabilistic_sharpe_ratio(
        selected_annualized_sharpe,
        benchmark,
        observations,
        skewness,
        raw_kurtosis,
        periods_per_year,
    )


def _self_test() -> None:
    se = iid_annualized_sharpe_se(1.5, 252, 252)
    assert math.isclose(se, 1.0022296571715914, rel_tol=1e-12)
    psr = probabilistic_sharpe_ratio(1.5, 0.0, 252, 0.0, 3.0, 252)
    assert math.isclose(psr, 0.9323717249467356, rel_tol=1e-12)
    monthly_equivalent = 1.5 * math.sqrt(12 / 252)
    monthly_psr = probabilistic_sharpe_ratio(
        monthly_equivalent, 0.0, 252, 0.0, 3.0, 12
    )
    assert math.isclose(psr, monthly_psr, rel_tol=1e-12)

    assert expected_maximum_sharpe_under_null([1.5]) == 0.0
    assert expected_maximum_sharpe_under_null([1.0, 1.0, 1.0]) == 0.0
    trials = [0.2, 0.4, 0.8, 1.1, 1.5]
    benchmark = expected_maximum_sharpe_under_null(trials)
    assert math.isclose(benchmark, 0.6254015702694767, rel_tol=1e-12)
    dsr = deflated_sharpe_ratio(1.5, trials, 252, 0.0, 3.0, 252)
    assert 0.0 <= dsr <= psr <= 1.0

    single_trial_dsr = deflated_sharpe_ratio(1.5, [1.5], 252, 0.0, 3.0, 252)
    assert math.isclose(single_trial_dsr, psr, rel_tol=1e-12)
    for invalid_call in (
        lambda: iid_annualized_sharpe_se(1.0, 1, 252),
        lambda: probabilistic_sharpe_ratio(1.0, 0.0, 1, 0.0, 3.0, 252),
        lambda: expected_maximum_sharpe_under_null([]),
        lambda: deflated_sharpe_ratio(1.0, [1.0, 2.0], 252, 0.0, 3.0, 252),
    ):
        try:
            invalid_call()
        except ValueError:
            pass
        else:
            raise AssertionError("invalid input did not raise ValueError")
    print(f"SE={se:.3f} PSR={psr:.3f} expected_max={benchmark:.3f} DSR={dsr:.3f}")


if __name__ == "__main__":
    _self_test()
