"""
Load Alpaca 1-minute RTH parquet and resample to session-aligned intraday bars.

Data directory default: ``SP500_1Min_Parquet_RTH_FULL`` (naive UTC → US/Eastern).
"""

from __future__ import annotations

import glob
import os
from pathlib import Path

import numpy as np
import pandas as pd

_REPO = Path(__file__).resolve().parents[2]
DEFAULT_ALPACA_RTH_DIR = _REPO / "SP500_1Min_Parquet_RTH_FULL"

# 09:30–16:00 ET = 390 minutes → 78 five-minute bars per session.
RTH_BARS_PER_SESSION_5MIN = 78
TRADING_DAYS_PER_YEAR = 252


def bars_per_year(*, bar_minutes: int = 5, sessions_per_year: int = TRADING_DAYS_PER_YEAR) -> int:
    if bar_minutes <= 0:
        raise ValueError("bar_minutes must be > 0")
    mins_per_session = 390
    per_session = mins_per_session // bar_minutes
    return int(per_session * sessions_per_year)


def load_minute_parquet(
    path: str | Path,
    *,
    naive_timestamp_tz: str = "UTC",
) -> pd.DataFrame:
    """Reuse breakout script loader (same parquet schema)."""
    from nasdaq_ma_breakout_daytrade import load_minute_parquet as _load

    return _load(str(path), naive_timestamp_tz=naive_timestamp_tz)


def resample_session_bars(minute: pd.DataFrame, bar_minutes: int = 5) -> pd.DataFrame:
    """RTH bars aligned to 09:30 open (offset 30min for 5m/60m grids)."""
    if bar_minutes <= 0:
        raise ValueError("bar_minutes must be > 0")
    offset = "30min" if bar_minutes in (5, 60) else None
    rule = f"{bar_minutes}min"
    if offset:
        agg = minute.resample(rule, offset=offset).agg(
            {
                "open": "first",
                "high": "max",
                "low": "min",
                "close": "last",
                "volume": "sum",
            }
        )
    else:
        agg = minute.resample(rule).agg(
            {
                "open": "first",
                "high": "max",
                "low": "min",
                "close": "last",
                "volume": "sum",
            }
        )
    out = agg.dropna(subset=["open", "close"]).sort_index()
    out["ret"] = out["close"].astype(np.float64).pct_change().fillna(0.0)
    return out


def daily_bars_from_minute(minute: pd.DataFrame) -> pd.DataFrame:
    daily = minute.resample("1D").agg(
        {
            "open": "first",
            "high": "max",
            "low": "min",
            "close": "last",
            "volume": "sum",
        }
    )
    daily = daily.dropna(subset=["close"]).sort_index()
    daily["ret"] = daily["close"].astype(np.float64).pct_change().fillna(0.0)
    return daily


def ohlcv_frame(bars: pd.DataFrame) -> pd.DataFrame:
    """Frame for intraday engines (DatetimeIndex + OHLCV + ret)."""
    need = ("open", "high", "low", "close", "ret")
    missing = [c for c in need if c not in bars.columns]
    if missing:
        raise KeyError(f"bars missing columns: {missing}")
    cols = list(need)
    if "volume" in bars.columns:
        cols.append("volume")
    return bars[cols].copy()


def compound_intraday_to_daily(r: pd.Series) -> pd.Series:
    """Chain intraday returns to one return per session date."""
    r = r.astype(np.float64)
    idx = pd.to_datetime(r.index)
    if getattr(idx, "tz", None) is not None:
        idx = idx.tz_localize(None)
    grouped = pd.Series(r.values, index=idx.normalize())
    return grouped.groupby(level=0).apply(lambda x: float(np.prod(1.0 + x) - 1.0))


def list_parquet_symbols(data_dir: str | Path = DEFAULT_ALPACA_RTH_DIR) -> list[str]:
    pattern = os.path.join(str(data_dir), "*.parquet")
    files = sorted(glob.glob(pattern))
    return [os.path.splitext(os.path.basename(p))[0] for p in files]


def load_symbol_bars(
    symbol: str,
    *,
    data_dir: str | Path = DEFAULT_ALPACA_RTH_DIR,
    bar_minutes: int = 5,
    start: str | None = None,
    end: str | None = None,
) -> tuple[pd.DataFrame, pd.DataFrame]:
    """Return (intraday_bars, daily_bars) from 1m parquet."""
    path = Path(data_dir) / f"{symbol}.parquet"
    if not path.is_file():
        raise FileNotFoundError(path)
    minute = load_minute_parquet(path)
    if start:
        minute = minute.loc[start:]
    if end:
        minute = minute.loc[:end]
    if minute.empty:
        raise ValueError(f"{symbol}: no minute data in window")
    intra = resample_session_bars(minute, bar_minutes=bar_minutes)
    daily = daily_bars_from_minute(minute)
    return ohlcv_frame(intra), ohlcv_frame(daily)


def load_equity_panels(
    symbols: list[str],
    *,
    data_dir: str | Path = DEFAULT_ALPACA_RTH_DIR,
    bar_minutes: int = 5,
    start: str | None = None,
    end: str | None = None,
    warmup_sessions: int = 0,
    verbose: bool = True,
) -> tuple[dict[str, pd.DataFrame], dict[str, pd.DataFrame]]:
    """Build intraday and daily equity dicts (skips missing/empty symbols)."""
    load_start = start
    if start and warmup_sessions > 0:
        load_start = (
            pd.Timestamp(start) - pd.Timedelta(days=int(max(warmup_sessions * 2, 10)))
        ).strftime("%Y-%m-%d")

    intra: dict[str, pd.DataFrame] = {}
    daily: dict[str, pd.DataFrame] = {}
    for sym in symbols:
        try:
            ib, db = load_symbol_bars(
                sym,
                data_dir=data_dir,
                bar_minutes=bar_minutes,
                start=load_start,
                end=end,
            )
            # Keep full intraday history for MA warm-up; caller trims return window.
            if start:
                db = db.loc[pd.Timestamp(start) :]
            if len(ib) < 100 or len(db) < 20:
                continue
            intra[sym] = ib
            daily[sym] = db
        except (FileNotFoundError, ValueError):
            continue
        if verbose and len(intra) % 50 == 0:
            print(f"  loaded {len(intra)} symbols …", flush=True)
    return intra, daily
