"""Pure causal math for the D48-XSMOM-28D-V2 research strategies."""

from __future__ import annotations

import math
from dataclasses import dataclass
from typing import Mapping

import numpy as np
import pandas as pd
from pandas import DataFrame, Series

LOOKBACK_DAYS = 28
MIN_ELIGIBLE_UNIVERSE = 10
BUCKET_SIZE = 4
SIGNAL_TIMEFRAME = pd.Timedelta(minutes=5)

REASON_SNAPSHOT_UNAVAILABLE = "xsmom_snapshot_unavailable"
REASON_UNSCORED_PASSTHROUGH = "xsmom_unscored_passthrough"
REASON_TOP4 = "xsmom_top4"
REASON_NOT_TOP4 = "xsmom_not_top4"
REASON_BOTTOM4 = "xsmom_bottom4"
REASON_NOT_BOTTOM4 = "xsmom_not_bottom4"
REASON_NOOP = "xsmom_noop"


@dataclass(frozen=True)
class RankMatrices:
    """Daily XSMOM state indexed by the completed daily candle timestamp."""

    scores: DataFrame
    ranks: DataFrame
    eligible_counts: Series
    snapshot_valid: Series


def as_utc_timestamp(value: object) -> pd.Timestamp:
    stamp = pd.Timestamp(value)
    return stamp.tz_localize("UTC") if stamp.tz is None else stamp.tz_convert("UTC")


def first_full_listing_day(listing_time: object) -> pd.Timestamp:
    """First UTC 1d candle timestamp not made partial by intraday onboarding."""
    listing = as_utc_timestamp(listing_time)
    day = listing.normalize()
    return day if listing == day else day + pd.Timedelta(days=1)


def normalize_daily_close(values: Series) -> Series:
    """Return a unique, sorted UTC-midnight float close series."""
    result = Series(
        pd.to_numeric(values, errors="coerce").to_numpy(),
        index=pd.DatetimeIndex(pd.to_datetime(values.index, utc=True)),
        dtype="float64",
    ).sort_index(kind="stable")
    if result.index.has_duplicates:
        raise ValueError("daily close series contains duplicate timestamps")
    if any(stamp != stamp.normalize() for stamp in result.index):
        raise ValueError("daily close timestamps must be UTC midnight")
    return result


def daily_score_series(daily_close: Series, listing_time: object) -> Series:
    """Compute causal 28-calendar-day scores from 29 complete positive closes."""
    close = normalize_daily_close(daily_close)
    if close.empty:
        return close.rename("xsmom_score")

    calendar = pd.date_range(close.index.min(), close.index.max(), freq="1D", tz="UTC")
    close = close.reindex(calendar)
    valid = Series(
        np.isfinite(close.to_numpy(dtype=float)) & (close.to_numpy(dtype=float) > 0),
        index=calendar,
    )
    complete = valid.rolling(
        LOOKBACK_DAYS + 1,
        min_periods=LOOKBACK_DAYS + 1,
    ).sum().eq(LOOKBACK_DAYS + 1)
    first_score_day = first_full_listing_day(listing_time) + pd.Timedelta(
        days=LOOKBACK_DAYS
    )
    complete &= calendar >= first_score_day

    scores = np.log(close / close.shift(LOOKBACK_DAYS)).where(complete)
    scores.name = "xsmom_score"
    return scores


def build_rank_matrices(
    daily_closes: Mapping[str, Series],
    listing_times: Mapping[str, object],
) -> RankMatrices:
    """Build deterministic daily score and rank matrices for the frozen universe."""
    pairs = sorted(daily_closes)
    if pairs != sorted(listing_times):
        raise ValueError("daily close pairs and listing-time pairs must match")

    score_columns = {
        pair: daily_score_series(daily_closes[pair], listing_times[pair])
        for pair in pairs
    }
    scores = DataFrame(score_columns).sort_index()
    eligible_counts = scores.notna().sum(axis=1).astype("int64")
    # Columns are canonical-string sorted. method='first' therefore implements
    # canonical-string ascending as the exact-tie breaker.
    ranks = scores.rank(axis=1, method="first", ascending=False)
    snapshot_valid = eligible_counts >= MIN_ELIGIBLE_UNIVERSE
    return RankMatrices(
        scores=scores,
        ranks=ranks,
        eligible_counts=eligible_counts,
        snapshot_valid=snapshot_valid,
    )


def decision_for_side(
    *,
    side: str,
    snapshot_valid: bool,
    pair_scoreable: bool,
    rank: float | int | None,
    eligible_count: int,
) -> tuple[bool, str, str | None]:
    """Return V2 allow/reason/bucket with frozen cold-start pass-through."""
    if not snapshot_valid:
        return False, REASON_SNAPSHOT_UNAVAILABLE, None
    if not pair_scoreable:
        return True, REASON_UNSCORED_PASSTHROUGH, None
    if rank is None or not math.isfinite(float(rank)):
        raise ValueError("scoreable pair must have a finite rank")

    ordinal = int(rank)
    if ordinal <= BUCKET_SIZE:
        bucket = "top4"
    elif ordinal >= eligible_count - BUCKET_SIZE + 1:
        bucket = "bottom4"
    else:
        bucket = "middle"

    if side == "long":
        allowed = bucket == "top4"
        return allowed, REASON_TOP4 if allowed else REASON_NOT_TOP4, bucket
    if side == "short":
        allowed = bucket == "bottom4"
        return allowed, REASON_BOTTOM4 if allowed else REASON_NOT_BOTTOM4, bucket
    raise ValueError(f"unsupported side: {side}")


def attach_signal_time_rank_state(
    dataframe: DataFrame,
    *,
    pair: str,
    matrices: RankMatrices,
) -> DataFrame:
    """Attach the rank snapshot available when each completed 5m signal is usable."""
    result = dataframe.copy()
    dates = pd.to_datetime(result["date"], utc=True, errors="raise")
    eligible_times = dates + SIGNAL_TIMEFRAME
    last_complete_days = eligible_times.dt.normalize() - pd.Timedelta(days=1)
    lookup = pd.DatetimeIndex(last_complete_days)

    scores = matrices.scores[pair].reindex(lookup).to_numpy(dtype=float)
    ranks = matrices.ranks[pair].reindex(lookup).to_numpy(dtype=float)
    counts = (
        matrices.eligible_counts.reindex(lookup).fillna(0).to_numpy(dtype="int64")
    )
    snapshots = (
        matrices.snapshot_valid.reindex(lookup).fillna(False).to_numpy(dtype=bool)
    )

    result["xsmom_signal_eligible_time"] = eligible_times
    result["xsmom_last_complete_utc_day"] = last_complete_days
    result["xsmom_score"] = scores
    result["xsmom_rank"] = ranks
    result["xsmom_eligible_universe_count"] = counts
    result["xsmom_rank_snapshot_valid"] = snapshots
    result["xsmom_pair_rank_eligible"] = np.isfinite(scores)

    long_decisions = [
        decision_for_side(
            side="long",
            snapshot_valid=bool(snapshot),
            pair_scoreable=math.isfinite(score),
            rank=rank,
            eligible_count=int(count),
        )
        for snapshot, score, rank, count in zip(
            snapshots, scores, ranks, counts, strict=True
        )
    ]
    short_decisions = [
        decision_for_side(
            side="short",
            snapshot_valid=bool(snapshot),
            pair_scoreable=math.isfinite(score),
            rank=rank,
            eligible_count=int(count),
        )
        for snapshot, score, rank, count in zip(
            snapshots, scores, ranks, counts, strict=True
        )
    ]
    result["xsmom_long_allowed"] = [item[0] for item in long_decisions]
    result["xsmom_long_reason"] = [item[1] for item in long_decisions]
    result["xsmom_short_allowed"] = [item[0] for item in short_decisions]
    result["xsmom_short_reason"] = [item[1] for item in short_decisions]
    return result


def apply_v2_gate(dataframe: DataFrame, *, enforce_rank_gate: bool) -> DataFrame:
    """Apply candidate or no-op behavior after authentic D48 raw markers exist."""
    result = dataframe.copy()
    raw_long = (
        result["enter_long"].fillna(0).astype(bool)
        if "enter_long" in result
        else Series(False, index=result.index)
    )
    raw_short = (
        result["enter_short"].fillna(0).astype(bool)
        if "enter_short" in result
        else Series(False, index=result.index)
    )
    result["xsmom_raw_long"] = raw_long
    result["xsmom_raw_short"] = raw_short
    result["xsmom_applied_decision"] = None
    result["xsmom_applied_reason"] = None

    if not enforce_rank_gate:
        result.loc[raw_long | raw_short, "xsmom_applied_decision"] = "pass"
        result.loc[raw_long | raw_short, "xsmom_applied_reason"] = REASON_NOOP
        return result

    long_allowed = result["xsmom_long_allowed"].fillna(False).astype(bool)
    short_allowed = result["xsmom_short_allowed"].fillna(False).astype(bool)
    result.loc[raw_long, "xsmom_applied_decision"] = np.where(
        long_allowed[raw_long], "pass", "block"
    )
    result.loc[raw_long, "xsmom_applied_reason"] = result.loc[
        raw_long, "xsmom_long_reason"
    ]
    result.loc[raw_short, "xsmom_applied_decision"] = np.where(
        short_allowed[raw_short], "pass", "block"
    )
    result.loc[raw_short, "xsmom_applied_reason"] = result.loc[
        raw_short, "xsmom_short_reason"
    ]
    if "enter_long" in result:
        result.loc[raw_long & ~long_allowed, "enter_long"] = 0
    if "enter_short" in result:
        result.loc[raw_short & ~short_allowed, "enter_short"] = 0
    return result
