"""Pure multi-speed EWMAC calculations for the research-only 1h strategy."""

from __future__ import annotations

from math import sqrt

import pandas as pd
from pandas import DataFrame, Series


EWMAC_SPEEDS = (
    (8, 32, 5.30),
    (16, 64, 3.75),
    (32, 128, 2.65),
)


def multi_speed_ewmac(
    close: Series,
    *,
    vol_span_hours: int = 72,
    forecast_cap: float = 20.0,
) -> DataFrame:
    """Return volatility-normalised components and their equal-weight forecast.

    The raw EWMAC is the fast/slow EMA price difference divided by lagging
    exponentially weighted daily price-unit volatility.  Scaling constants are
    frozen per speed and each component is capped before combination.
    """
    if vol_span_hours < 2 or forecast_cap <= 0:
        raise ValueError("EWMAC volatility span and cap must be positive")
    prices = pd.to_numeric(close, errors="coerce").astype("float64")
    if (prices.dropna() <= 0).any():
        raise ValueError("EWMAC close prices must be positive")

    hourly_returns = prices.pct_change(fill_method=None)
    hourly_return_vol = hourly_returns.ewm(
        span=vol_span_hours,
        min_periods=vol_span_hours,
        adjust=False,
    ).std(bias=False)
    daily_price_vol = prices * hourly_return_vol * sqrt(24.0)
    daily_price_vol = daily_price_vol.where(daily_price_vol > 0)

    result = DataFrame(index=prices.index)
    components = []
    for fast, slow, scalar in EWMAC_SPEEDS:
        fast_ema = prices.ewm(span=fast, min_periods=fast, adjust=False).mean()
        slow_ema = prices.ewm(span=slow, min_periods=slow, adjust=False).mean()
        component = ((fast_ema - slow_ema) / daily_price_vol * scalar).clip(
            lower=-forecast_cap,
            upper=forecast_cap,
        )
        name = f"ewmac_{fast}_{slow}"
        result[name] = component
        components.append(name)

    result["ewmac_forecast"] = result[components].mean(axis=1, skipna=False)
    result["ewmac_daily_price_vol"] = daily_price_vol
    return result


def hysteresis_signals(
    forecast: Series,
    *,
    entry_threshold: float = 5.0,
    exit_threshold: float = 0.0,
) -> DataFrame:
    """Create causal threshold-cross entries and zero-axis exits."""
    if entry_threshold <= exit_threshold or exit_threshold < 0:
        raise ValueError("thresholds must satisfy entry > exit >= 0")
    values = pd.to_numeric(forecast, errors="coerce")
    previous = values.shift(1)
    return DataFrame(
        {
            "enter_long": (values >= entry_threshold) & (previous < entry_threshold),
            "enter_short": (values <= -entry_threshold) & (previous > -entry_threshold),
            "exit_long": (values <= exit_threshold) & (previous > exit_threshold),
            "exit_short": (values >= -exit_threshold) & (previous < -exit_threshold),
        },
        index=values.index,
    )
